| |
| """Reproduce the documented dev metric (first-token top-1 / top-5 / CE) from packaged files. |
| |
| Protocol (identical to the project's scripts/04_diag_checkpoint.py): |
| rows = first 1024 rows of data/dev/dev_ids.npy with dev_ids[k,1] == <|FORWARD|> |
| logits = model(ids)[k, mask_start-1] (the position right before the first tactic token) |
| target = ids[k, mask_start] |
| This is the ONLY comparable dev protocol: the numbers printed by the trainer use a |
| length-bucketed reshuffle of the dev split and must not be compared across datasets. |
| |
| python eval_dev.py # checkpoints/e6, 1024 rows |
| python eval_dev.py --ckpt checkpoints/stage3-e3 --rows 256 |
| python eval_dev.py --all --compare # every variant + compare with metrics.json |
| """ |
| import argparse |
| import json |
| import os |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
|
|
| from common import ROOT, VARIANTS, load_config, load_model, pick_device, specials |
|
|
| CFG = load_config() |
|
|
|
|
| def dev_eval(net, device, rows=1024, batch=64, verbose=False): |
| sp = specials() |
| pad = sp['<|pad|>'] |
| ids_all = np.load(os.path.join(ROOT, 'data/dev/dev_ids.npy'), mmap_mode='r') |
| ms_all = np.load(os.path.join(ROOT, 'data/dev/dev_mask_start.npy')) |
| ln_all = np.load(os.path.join(ROOT, 'data/dev/dev_len.npy')) |
| fwd = [k for k in range(len(ln_all)) if ids_all[k, 1] == sp['<|FORWARD|>']][:rows] |
| n1 = n5 = n = 0 |
| ce_sum = 0.0 |
| with torch.no_grad(): |
| for s in range(0, len(fwd), batch): |
| chunk = fwd[s:s + batch] |
| seqs = [np.asarray(ids_all[k, :int(ln_all[k])], dtype=np.int64) for k in chunk] |
| m = max(len(x) for x in seqs) |
| arr = np.full((len(chunk), m), pad, dtype=np.int64) |
| for j, x in enumerate(seqs): |
| arr[j, :len(x)] = x |
| t = torch.from_numpy(arr).to(device) |
| ms = torch.tensor([min(int(ms_all[k]), m - 1) for k in chunk], device=device) |
| ar = torch.arange(len(chunk), device=device) |
| lg = net(t)['logits'][ar, (ms - 1).clamp(min=0)] |
| tgt = t[ar, ms.clamp(max=m - 1)] |
| p5 = lg.topk(5, -1).indices |
| n1 += int((p5[:, 0] == tgt).sum()) |
| n5 += int((p5 == tgt[:, None]).any(1).sum()) |
| ce_sum += float(F.cross_entropy(lg.float(), tgt, reduction='sum')) |
| n += len(chunk) |
| if verbose: |
| print(f' ... {n}/{len(fwd)} rows', flush=True) |
| return {'rows': n, 'top1': n1 / n, 'top5': n5 / n, 'ce': ce_sum / n} |
|
|
|
|
| def documented(variant): |
| m = json.load(open(os.path.join(ROOT, 'metrics.json')))['dev'] |
| key = {'checkpoints/e6': 'stage4_e6_full_best', |
| 'checkpoints/stage3-e3': 'stage3_e3_best'}.get(variant) |
| return m.get(key) if key else None |
|
|
|
|
| if __name__ == '__main__': |
| ap = argparse.ArgumentParser() |
| ap.add_argument('--ckpt', default=VARIANTS[0]) |
| ap.add_argument('--all', action='store_true', help='evaluate every packaged variant') |
| ap.add_argument('--rows', type=int, default=1024) |
| ap.add_argument('--batch', type=int, default=64) |
| ap.add_argument('--device', default=None) |
| ap.add_argument('--compare', action='store_true', help='compare with metrics.json') |
| a = ap.parse_args() |
| device = pick_device(a.device) |
| todo = VARIANTS if a.all else (a.ckpt,) |
| worst = 0.0 |
| for ck in todo: |
| net, _cfg, device = load_model(ck, device) |
| r = dev_eval(net, device, rows=a.rows, batch=a.batch) |
| line = (f'{ck:22s} rows={r["rows"]:4d} top1={r["top1"]:.4f} ' |
| f'top5={r["top5"]:.4f} ce={r["ce"]:.4f}') |
| doc = documented(ck) |
| if a.compare and doc and r['rows'] == 1024: |
| d = max(abs(r['top1'] - doc['top1']), abs(r['top5'] - doc['top5']), |
| abs(r['ce'] - doc['ce'])) |
| worst = max(worst, d) |
| line += f' | 文档值 {doc["top1"]:.4f}/{doc["top5"]:.4f}/{doc["ce"]:.4f} 偏差 {d:.5f}' |
| line += ' ✅' if d < 5e-4 else ' ❌' |
| print(line, flush=True) |
| if a.compare: |
| print(f'\n最大偏差 {worst:.5f} → ' + ('PASS' if worst < 5e-4 else 'FAIL')) |
|
|