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())