"""Per-position accuracy of the seq draft — where does the chain die? python3 eval_seq_pos.py [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}")