Download code/eval_seq_pos.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 1.57 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/eval_seq_pos.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/eval_seq_pos.py
-
curl -L -o eval_seq_pos.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/eval_seq_pos.py
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}") | |