spec100m / code /eval_streak.py
Akahsizrr's picture
Upload code/eval_streak.py with huggingface_hub
afcfe02 verified
Raw History Blame Contribute Delete
1.05 kB
"""Heavier eval of v6 spec heads: multiple windows x more anchors."""
import sys
import torch
sys.path.insert(0, "/root/spec100m")
from config import Config
from model import build_model
from data_pipeline import load_tokens
from train_medusa import eval_streak
device = "cuda"
cfg = Config.v5_500m()
model = build_model(cfg, device)
data = load_tokens("mixture500m").to(device)
for ckpt_path in sys.argv[1:]:
ckpt = torch.load(ckpt_path, map_location=device)
missing, unexpected = model.load_state_dict(ckpt["model"], strict=False)
model.eval()
streaks, accs = [], []
for _ in range(8): # 8 windows x 64 anchors = 512 samples
s, a = eval_streak(model, cfg, data, seq=1024, n_anchors=64, device=device)
streaks.append(s)
accs.append(a)
sm = sum(streaks) / len(streaks)
am = sum(accs) / len(accs)
print(f"{ckpt_path} (step {ckpt.get('step')}): "
f"streak {sm:.2f} (min {min(streaks):.2f} max {max(streaks):.2f}) "
f"head-0 acc {am:.2f}")