""" Standalone usage example -- everything needed to run this model is in this same folder (model.py, model.safetensors, tokenizer.json/vocab.json/merges.txt, config.json). Requires: torch, tokenizers, safetensors pip install torch tokenizers safetensors Usage: python example.py """ import json import torch from tokenizers import Tokenizer from safetensors.torch import load_model from model import ChatGPTMini, ModelConfig # ---- load config + tokenizer ---- with open("config.json") as f: cfg_dict = json.load(f) tokenizer = Tokenizer.from_file("tokenizer.json") cfg = ModelConfig( vocab_size=cfg_dict["vocab_size"], d_model=cfg_dict["d_model"], n_layer=cfg_dict["n_layer"], n_head=cfg_dict["n_head"], d_ff=cfg_dict["d_ff"], max_seq_len=cfg_dict["max_seq_len"], dropout=0.0, # no dropout at inference time pad_token_id=cfg_dict["pad_token_id"], ) # ---- load model weights ---- model = ChatGPTMini(cfg) load_model(model, "model.safetensors") model.eval() # ---- generate ---- bos_id = tokenizer.token_to_id("") eos_id = tokenizer.token_to_id("") chat_id = tokenizer.token_to_id("") superchat_id = tokenizer.token_to_id("") print("== chat mode ==") prompt = torch.tensor([[bos_id, chat_id]], dtype=torch.long) out = model.generate(prompt, max_new_tokens=30, temperature=0.9, top_k=40, top_p=0.9, eos_token_id=eos_id) print(tokenizer.decode(out[0].tolist(), skip_special_tokens=True)) print("\n== superchat mode ==") prompt = torch.tensor([[bos_id, superchat_id]], dtype=torch.long) out = model.generate(prompt, max_new_tokens=40, temperature=0.8, top_k=40, top_p=0.9, eos_token_id=eos_id) print(tokenizer.decode(out[0].tolist(), skip_special_tokens=True)) print("\n== batched generation (fast, recommended for real use) ==") prompt = torch.tensor([[bos_id, chat_id]] * 8, dtype=torch.long) out = model.generate(prompt, max_new_tokens=20, temperature=0.9, top_k=40, top_p=0.9, eos_token_id=eos_id) for row in out: print(" -", tokenizer.decode(row.tolist(), skip_special_tokens=True))