| """ |
| 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 |
|
|
| |
| 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, |
| pad_token_id=cfg_dict["pad_token_id"], |
| ) |
|
|
| |
| model = ChatGPTMini(cfg) |
| load_model(model, "model.safetensors") |
| model.eval() |
|
|
| |
| 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)) |
|
|