File size: 4,769 Bytes
32c0c6c | 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 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | #!/usr/bin/env python
"""One-command verification of this release package.
python verify_package.py # full check (1024 dev rows, both variants)
python verify_package.py --quick # 128 dev rows (fast, on CPU)
Checks, in order:
1. architecture in config.json rebuilds a model with the documented parameter count
2. every packaged variant reproduces the documented dev protocol numbers
(top-1 / top-5 / CE on 1024 fixed rows, tolerance 5e-4)
3. a real example state produces tactics (inference wiring works end to end)
4. SHA256SUMS verifies the big files byte-for-byte
Exit code 0 only if every check passes.
"""
import argparse
import hashlib
import json
import os
import sys
from common import ROOT, VARIANTS, load_config, load_model, load_tokenizer, pick_device, \
encode_state, whitelist, specials
from eval_dev import dev_eval, documented
def sha256(path, chunk=1 << 22):
h = hashlib.sha256()
with open(path, 'rb') as fh:
while True:
b = fh.read(chunk)
if not b:
break
h.update(b)
return h.hexdigest()
def check_sums():
p = os.path.join(ROOT, 'SHA256SUMS')
if not os.path.exists(p):
return 'SHA256SUMS 不存在(跳过)', False
bad, n = [], 0
for line in open(p):
line = line.strip()
if not line or line.startswith('#'):
continue
want, name = line.split(None, 1)
name = name.lstrip('*')
f = os.path.join(ROOT, name)
if not os.path.exists(f):
bad.append(f'{name} 缺失')
continue
n += 1
if sha256(f) != want:
bad.append(f'{name} 校验和不符')
return ('SHA256SUMS:%d 个文件全部匹配' % n if not bad else ';'.join(bad)), not bad
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--quick', action='store_true', help='128 dev rows instead of 1024')
ap.add_argument('--device', default=None)
ap.add_argument('--skip-sums', action='store_true')
a = ap.parse_args()
rows = 128 if a.quick else 1024
device = pick_device(a.device)
cfg = load_config()
ok_all = True
print(f'package : {ROOT}')
print(f'device : {device} dev rows: {rows}\n')
print('[1] 参数数量')
want = cfg['n_params']
for ck in VARIANTS:
net, _, _ = load_model(ck, device)
got = sum(p.numel() for p in net.parameters())
good = got == want
ok_all &= good
print(f' {ck:22s} {got:,} (config: {want:,}) {"OK" if good else "MISMATCH"}')
del net
print('\n[2] dev 协议复算(top-1 / top-5 / CE)')
for ck in VARIANTS:
net, _, device = load_model(ck, device)
r = dev_eval(net, device, rows=rows)
doc = documented(ck)
line = f' {ck:22s} top1={r["top1"]:.4f} top5={r["top5"]:.4f} ce={r["ce"]:.4f}'
if doc and rows == 1024:
d = max(abs(r['top1'] - doc['top1']), abs(r['top5'] - doc['top5']),
abs(r['ce'] - doc['ce']))
good = d < 5e-4
ok_all &= good
line += f' | documented {doc["top1"]:.4f}/{doc["top5"]:.4f}/{doc["ce"]:.4f} Δ={d:.5f} {"OK" if good else "MISMATCH"}'
elif doc:
line += f' | documented {doc["top1"]:.4f}/{doc["top5"]:.4f}/{doc["ce"]:.4f} (只用 1024 行才判定)'
print(line)
del net
print('\n[3] 推理连通性(示例样本 0)')
net, _, device = load_model(VARIANTS[0], device)
tok = load_tokenizer()
rec = [json.loads(l) for l in open(os.path.join(ROOT, 'examples/dev_sample.jsonl'))][0]
import torch
import torch.nn.functional as F
p = encode_state(tok, rec['state'], cfg)
ids = torch.tensor([p], device=device)
with torch.no_grad():
lg = net(ids)['logits'][0, -1]
allow = torch.full_like(lg, float('-inf'))
allow[torch.tensor(whitelist(), device=device)] = 0.0
top = int((lg + allow).argmax())
sp = specials()
pred = tok.decode([top])
truth = rec['true_first_token']
hit = bool(pred.strip() == truth.strip())
print(f' 状态: {rec["state"].strip().splitlines()[-1][:60]}')
print(f' 模型首个 token: {pred!r} | 该真实状态的 tactic 首 token: {truth!r} '
f'{"命中" if hit else "未命中(正常:模型只有 0.31 的首 token 命中率,不影响本包可用性)"}')
del net
if not a.skip_sums:
print('\n[4] 校验和')
msg, good = check_sums()
ok_all &= good
print(f' {msg}')
print('\n=== ' + ('PASS:本包自洽,文档中的数字可复现' if ok_all else 'FAIL:见上面标记'))
return 0 if ok_all else 1
if __name__ == '__main__':
sys.exit(main())
|