NanoChat-28M-V1 / example.py
Batman55072's picture
Initial upload
6e23275
Raw
History Blame Contribute Delete
2.08 kB
"""
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("<bos>")
eos_id = tokenizer.token_to_id("<eos>")
chat_id = tokenizer.token_to_id("<chat>")
superchat_id = tokenizer.token_to_id("<superchat>")
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))