queen-v2 / chat.py
Crawlnet's picture
queen-v2 (4ba917bb5096)
f306176
Raw History Blame Contribute Delete
3.25 kB
"""
Chat with a released Crawlnet queen from its folder. Copied into every release as chat.py.
pip install torch tiktoken safetensors
git clone https://github.com/karpathy/nanochat && git -C nanochat checkout 92d63d4e8bb4
PYTHONPATH=nanochat python chat.py "What is staking?" [--dir .] [--temperature 0.7] [--top-k 50] [--max-tokens 200]
Needs only nanochat (MIT) for the model and tokenizer classes; no Crawlnet code.
"""
import argparse
import base64
import json
import os
import tiktoken
import torch
from safetensors.torch import load_file
from nanochat.engine import Engine
from nanochat.gpt import GPT, GPTConfig
from nanochat.tokenizer import RustBPETokenizer
def load(folder, device):
with open(os.path.join(folder, "config.json")) as f:
cfg = json.load(f)
with open(os.path.join(folder, "tokenizer.json")) as f:
tj = json.load(f)
ranks = {}
with open(os.path.join(folder, "tokenizer.tiktoken")) as f:
for line in f:
tok, rank = line.split()
ranks[base64.b64decode(tok)] = int(rank)
enc = tiktoken.Encoding(name="queen", pat_str=tj["pattern"], mergeable_ranks=ranks, special_tokens=tj["special_tokens"])
tokenizer = RustBPETokenizer(enc, tj["bos_token"])
with torch.device("meta"):
model = GPT(GPTConfig(**cfg["model_config"]))
model.to_empty(device=device)
model.init_weights() # rotary buffers; the weights are overwritten below
state = load_file(os.path.join(folder, "model.safetensors"), device="cpu")
if device != "cuda": # bf16 weights run in fp32 on CPU
state = {k: (v.float() if v.is_floating_point() else v) for k, v in state.items()}
model.load_state_dict({k: v.to(device) for k, v in state.items()}, strict=True, assign=True)
model.eval()
return model, tokenizer
def answer(model, tokenizer, question, max_tokens=200, temperature=0.7, top_k=50, seed=42):
sp = tokenizer.encode_special
bos = tokenizer.get_bos_token_id()
prompt = [bos, sp("<|user_start|>")] + tokenizer.encode(question) + [sp("<|user_end|>"), sp("<|assistant_start|>")]
stop = {bos, sp("<|assistant_end|>"), sp("<|user_start|>")}
out = []
for column, _ in Engine(model, tokenizer).generate(prompt, num_samples=1, max_tokens=max_tokens,
temperature=temperature, top_k=top_k, seed=seed):
if column[0] in stop:
break
out.append(column[0])
return tokenizer.decode(out).strip()
def main():
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("question")
ap.add_argument("--dir", default=os.path.dirname(os.path.abspath(__file__)))
ap.add_argument("--max-tokens", type=int, default=200)
ap.add_argument("--temperature", type=float, default=0.7)
ap.add_argument("--top-k", type=int, default=50)
ap.add_argument("--seed", type=int, default=42)
args = ap.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
model, tokenizer = load(args.dir, device)
print(answer(model, tokenizer, args.question, args.max_tokens, args.temperature, args.top_k, args.seed))
if __name__ == "__main__":
main()