delta-slotstack-236m / margin_test.py
Aurov's picture
initial upload
bbb9a9b verified
Raw
History Blame Contribute Delete
7.38 kB
#!/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()