File size: 7,383 Bytes
bbb9a9b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
#!/usr/bin/env python3
"""
================================================================================
margin_test.py — self 到底是几个零? + 画板还剩多少余量
================================================================================

一. 为什么日志里的 0.000 可能是【精确的 0.0】
---------------------------------------------
cross_entropy 的实现是 logsumexp − logit_target。当目标 logit 领先 Δ 时
    loss ≈ Σ_{j≠t} exp(−Δ_j)
而 fp32 里 1+x 在 x < eps/2 = 5.96e−8 时直接舍入成 1, log(1)=0。
实测: 领先 20 以上, fp32 的 cross_entropy 就返回精确 0.0。
本架构 logit 值域 ≥ ±√d·RMS(γ) = ±36, 能拉开 70+, 远超这个门槛。

★ 本脚本用 float64 + log1p 绕过这个地板:
      s    = Σ_{j≠t} exp(lg_j − lg_t)     ← float64
      loss = log1p(s)                      ← log1p 对极小值精确
  exp(−700) 仍在 fp64 范围内, 所以能一路量到 1e−300 量级。
  这才是"小数点后几个零"的真实答案。

二. 为什么该改看 margin
-----------------------
    margin = logit(目标) − max_{j≠目标} logit(j)
loss 一旦打满就【精确为 0】, 于是有一个巨大的死区: 假设当前 margin=70,
你往某个头里塞东西把它打掉 45, loss 仍然是 0.0 —— 仪表纹丝不动。
把"改一点就立刻偏离"的探测器变成了"撞到底才动"的开关, 与设计意图相反。

margin 连续、无死区、单位就是 logit, 塞进去多少就掉多少, 1:1 可读。
它直接回答: 【这块画板还剩多少余量可以被涂掉】。

三. 用法
--------
    python margin_test.py --model dual_d576_L12.py \\
        --ckpt ./dual_out/latest.pt --batches 8
    python margin_test.py --model delta_d1280_h10_l10.py \\
        --ckpt ./d1280_out/latest.pt
    多臂文件会自动把每一臂都测一遍并列出来 (std vs dual, base vs pyramid …)
================================================================================
"""
import argparse, importlib.util, inspect, math
import numpy as np
import torch
import torch.nn.functional as F


def load_module(path):
    spec = importlib.util.spec_from_file_location("expmod", path)
    m = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(m)
    return m


def build_models(m, ck):
    """兼容单臂(DeltaLM / ckpt['model'])与多臂(LM(kind) / BUILDERS / ckpt['models'])。"""
    out = {}
    if ck is not None and "models" in ck:              # 多臂
        for kind, sd in ck["models"].items():
            if hasattr(m, "BUILDERS"):
                mod = m.BUILDERS[kind]()
            else:
                mod = m.LM(kind)
            mod.load_state_dict(sd)
            out[kind] = mod
        return out
    # 单臂
    for name in ("DeltaLM", "LM", "Model", "GPT"):
        if hasattr(m, name):
            cls = getattr(m, name)
            try:
                mod = cls() if len(inspect.signature(cls).parameters) == 0 else cls("std")
            except Exception:
                continue
            if ck is not None:
                mod.load_state_dict(ck["model"] if "model" in ck else ck)
            out["model"] = mod
            return out
    raise RuntimeError("认不出模型类, 请手工指定")


def probe_tokens(m):
    if hasattr(m, "_PROBE_TOK"):
        return dict(m._PROBE_TOK)
    d = {}
    for nm, at in (("self", "TOK_SELF"), ("prev", "TOK_PREV"), ("first", "TOK_FIRST"),
                   ("dup", "TOK_DUP"), ("b16", "TOK_B16"),
                   ("b128", "TOK_B128"), ("b512", "TOK_B512")):
        if hasattr(m, at):
            d[nm] = getattr(m, at)
    return d


def make_loader(m, insert_p):
    if getattr(m, "SMOKE", False):
        return m.SyntheticLoader(m.MICRO_BATCH, m.SEQ_LEN, 3)
    sig = inspect.signature(m.ShardLoader.__init__).parameters
    kw = {}
    if "insert_p" in sig:
        kw["insert_p"] = insert_p
    if "seed" in sig:
        kw["seed"] = getattr(m, "SEED", 42) + 999
    return m.ShardLoader(m.DATA_DIR, "val", m.MICRO_BATCH, m.SEQ_LEN, **kw)


def get_logits(model, x):
    out = model(x)
    return out[0] if isinstance(out, (tuple, list)) else out


@torch.no_grad()
def measure(model, batches, tok, device):
    """返回 (fp64精确loss数组, margin数组, fp32报告loss数组)。"""
    L64, MG, L32 = [], [], []
    for x, y in batches:
        msk = (x == tok)
        if not msk.any():
            continue
        lg = get_logits(model, x)[msk].double()          # (n, V)
        t = y[msk][:, None]
        lt = lg.gather(1, t)                              # 目标 logit
        rest = lg.scatter(1, t, float("-inf"))
        # fp64 + log1p: s = Σ_{j≠t} exp(lg_j − lg_t), 可量到 1e−300
        s = (rest - lt).exp().sum(1)
        L64.append(torch.log1p(s).cpu())
        MG.append((lt.squeeze(1) - rest.max(1).values).cpu())
        L32.append(F.cross_entropy(lg.float(), y[msk], reduction="none").cpu())
    if not L64:
        return None
    return (torch.cat(L64).numpy(), torch.cat(MG).numpy(), torch.cat(L32).numpy())


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--ckpt", default=None)
    ap.add_argument("--batches", type=int, default=8)
    ap.add_argument("--insert-p", type=float, default=0.05)
    a = ap.parse_args()

    m = load_module(a.model)
    dev = m.device
    ck = None
    if a.ckpt:
        ck = torch.load(a.ckpt, map_location=dev, weights_only=False)
    models = {k: v.to(dev).eval() for k, v in build_models(m, ck).items()}
    toks = probe_tokens(m)

    print("=" * 104)
    print(f"  探针精确度测量   {a.model}   {a.ckpt or '随机初始化(无结论)'}"
          + (f"   step={ck.get('step','?')}" if ck else ""))
    print(f"  臂: {', '.join(models)}   探针: {', '.join(toks)}")
    print("=" * 104)
    print(f"\n  fp32 地板: 目标 logit 领先 ~20 以上, cross_entropy 就返回【精确 0.0】")
    print(f"  本脚本用 float64 + log1p 绕过, 可量到 1e-300 量级\n")

    ld = make_loader(m, a.insert_p)
    ld.reset()
    batches = [ld.next_batch(dev) for _ in range(a.batches)]

    for kind, mod in models.items():
        print(f"  ── 臂: {kind} " + "─" * (92 - len(kind)))
        print(f"  {'探针':<8}{'fp32报告':<12}{'fp64真值(中位)':<18}"
              f"{'margin 中位':<14}{'margin 最小':<14}{'fp32为精确0的比例'}")
        for nm, tok in toks.items():
            r = measure(mod, batches, tok, dev)
            if r is None:
                print(f"  {nm:<8}无样本")
                continue
            l64, mg, l32 = r
            z = float((l32 == 0.0).mean())
            print(f"  {nm:<8}{np.median(l32):<12.4f}{np.median(l64):<18.3e}"
                  f"{np.median(mg):<14.2f}{mg.min():<14.2f}{z*100:.1f}%")
        print()

    print("  ★ fp32报告=0.0000 而 fp64真值=1e-30 量级 ⇒ 不是'很多个零', 是浮点装不下。")
    print("  ★ margin 才是画板刻度: 它是【还能被涂掉多少 logit】。margin 中位 70 意味着")
    print("    你往头里塞东西打掉 45, loss 仍精确为 0 —— 死区极大, 别拿 loss 当探测器。")
    print("  ★ 多臂对照直接看 margin 那两列: 谁的余量大, 谁的画板就更能被写进去。")
    print("=" * 104)


if __name__ == "__main__":
    main()