Download code/eval_dev.py from Duoia/duotactic: direct link, hf CLI and curl.
- Browser
- Download file 4.15 kB
-
https://huggingface.co/Duoia/duotactic/resolve/main/code/eval_dev.py
- Command line
-
hf download hf://Duoia/duotactic/code/eval_dev.py
-
curl -L -o eval_dev.py https://huggingface.co/Duoia/duotactic/resolve/main/code/eval_dev.py
4.15 kB
| #!/usr/bin/env python | |
| """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')) | |