duotactic / code /verify_package.py
Duoia's picture
duotactic full package: checkpoints, tokenizer, config, code, docs
32c0c6c verified
Raw
History Blame Contribute Delete
4.77 kB
#!/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())