bear-240m-cpt / inference.py
Dummy9898's picture
Release Mesosfer Bear AI Model checkpoint (bear_cpt)
b2140f4 verified
Raw
History Blame Contribute Delete
2.63 kB
"""
Standalone Inference Runner for Mesosfer Bear AI
"""
import os
import sys
import json
import argparse
import torch
from engine.transformer import BearTransformer, BearConfig
from engine.tokenizer import BearTokenizer
def main():
parser = argparse.ArgumentParser(description="Mesosfer Bear AI Standalone Inference")
parser.add_argument("--prompt", type=str, default="Halo, jelaskan apa itu kecerdasan buatan dalam 2 kalimat.")
parser.add_argument("--checkpoint", type=str, default="bear_model.pt")
parser.add_argument("--config", type=str, default="config.json")
parser.add_argument("--tokenizer", type=str, default="bear_tokenizer.json")
parser.add_argument("--max-tokens", type=int, default=256)
parser.add_argument("--temperature", type=float, default=0.7)
parser.add_argument("--top-p", type=float, default=0.9)
parser.add_argument("--top-k", type=int, default=40)
parser.add_argument("--thinking", action="store_true", help="Enable XTML thinking mode")
args = parser.parse_args()
device = "cuda" if torch.cuda.is_available() else ("mps" if hasattr(torch.backends, "mps") and torch.backends.mps.is_available() else "cpu")
print(f"Loading Bear AI model on {device}...")
# 1. Load Tokenizer
tokenizer = BearTokenizer.load(args.tokenizer) if os.path.exists(args.tokenizer) else BearTokenizer()
# 2. Load Config & Model
with open(args.config, "r", encoding="utf-8") as f:
cfg_dict = json.load(f)
config = BearConfig.from_dict(cfg_dict)
model = BearTransformer(config)
ckpt = torch.load(args.checkpoint, map_location=device, weights_only=False)
state_dict = ckpt.get("model_state", ckpt)
model.load_state_dict(state_dict)
model.to(device)
model.eval()
# 3. Format Prompt
conv = [
{"role": "system", "content": "Anda adalah asisten AI Bear yang cerdas, ringkas, dan ramah."},
{"role": "user", "content": args.prompt}
]
formatted = tokenizer.apply_chat_template(conv, thinking=args.thinking)
input_ids = torch.tensor([tokenizer.encode(formatted)], dtype=torch.long, device=device)
print(f"\nPrompt: {args.prompt}\n" + "=" * 60)
print("Generating response (streaming):\n")
# 4. Generate
with torch.no_grad():
out = model.generate(
input_ids,
max_new_tokens=args.max_tokens,
temperature=args.temperature,
top_p=args.top_p,
top_k=args.top_k,
)
generated_text = tokenizer.decode(out[0].tolist())
print(generated_text)
print("=" * 60)
if __name__ == "__main__":
main()