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