| |
| """ |
| ================================================================================ |
| 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() |
| t = y[msk][:, None] |
| lt = lg.gather(1, t) |
| rest = lg.scatter(1, t, float("-inf")) |
| |
| 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() |
|
|