spec100m / code /eval_seq_pos.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
1.57 kB
"""Per-position accuracy of the seq draft — where does the chain die?
python3 eval_seq_pos.py <ckpt> [chain]"""
import sys, torch
from config import Config
from model import build_model
from data_pipeline import load_tokens
ckpt_path = sys.argv[1]
chain = int(sys.argv[2]) if len(sys.argv) > 2 else 32
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)
print(f"loaded {ckpt_path} (step {ckpt.get('step')}) "
f"missing={len(msd)-len(sd)}")
model.eval()
data = load_tokens("mixture500m").to("cuda")
E = model.tok_emb.weight
seq = 1024
start = torch.randint(0, data.numel() - seq - 1, (1,), device="cuda")
ids = data[start:start + seq].unsqueeze(0)
with torch.no_grad():
h = model(ids)
base_pred = model.lm_head(h).argmax(-1)[0]
G = cfg.medusa_cond_group
anchors = torch.randint(G, seq - chain - 3, (256,), device="cuda")
h_a = h[0, anchors]
pend = (h_a @ E.T).argmax(-1, keepdim=True)
hist = torch.stack([ids[0, t - G + 1:t + 1] for t in anchors])
draft = model.spec_draft_seq(h_a, prefix_ids=pend, hist_ids=hist, n=chain)
tgt = torch.stack([base_pred[t + 1:t + 1 + chain] for t in anchors])
match = (draft == tgt).float().mean(0)
streak = match.cumprod(0)
print("pos: acc cumstreak-contrib")
for i in range(chain):
print(f" +{i+2:2d}: {match[i].item():.3f} {streak[i].item():.3f}")
print(f"\nmean streak ~{streak[-1].item():.2f}")