spec100m / code /eval_seq_text.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
1.7 kB
"""Sampled seq-draft text quality: emit a draft chain, decode it.
python3 eval_seq_text.py <ckpt> [prompt]"""
import sys, torch
from config import Config
from model import build_model
from data_pipeline import load_tokens
import tiktoken
enc = tiktoken.get_encoding("gpt2")
ckpt_path = sys.argv[1]
prompt = sys.argv[2] if len(sys.argv) > 2 else "The history of the Roman Empire"
cfg = Config.v5_500m()
model = build_model(cfg, "cuda")
ckpt = torch.load(ckpt_path, map_location="cuda")
msd = model.state_dict()
sd = {k: v for k, v in ckpt["model"].items()
if k in msd and msd[k].shape == v.shape}
model.load_state_dict(sd, strict=False)
model.eval()
G = cfg.medusa_cond_group
ids = torch.tensor([enc.encode(prompt)], device="cuda")
with torch.no_grad():
h = model(ids)
h_a = h[:, -1] # [1,d]
pend = (h_a @ model.tok_emb.weight.T).argmax(-1, keepdim=True)
hist = ids[:, -G:] # last G committed toks
print(f"base argmax next: {enc.decode(pend[0].tolist())!r}")
# argmax chain
d1 = model.spec_draft_seq(h_a, prefix_ids=pend, hist_ids=hist, n=32)
print(f"\nARGMAX draft (32): {enc.decode(d1[0].tolist())!r}")
# sampled chains
for i in range(3):
d2 = model.spec_draft_seq(h_a, prefix_ids=pend, hist_ids=hist, n=32,
sample=True, temperature=0.9, top_p=0.9)
print(f"SAMPLED draft {i}: {enc.decode(d2[0].tolist())!r}")
# base reference: sample from base at same anchor
lg = (h_a @ model.tok_emb.weight.T).float() / 0.9
b_tok = torch.multinomial(lg.softmax(-1), 1)
print(f"\nbase sample next: {enc.decode(b_tok[0].tolist())!r}")