#!/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()