File size: 1,573 Bytes
a8f07a3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
"""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}")