#!/usr/bin/env python3 """ ================================================================================ head_test.py — 头独立性检验 + 单头消融归因 ================================================================================ 这个脚本回答的是"可解释性"本身的问题, 不是"跟加性残差比怎么样"。 A. 头子空间独立性 (纯看权重, 秒级, 不用数据) ------------------------------------------------------------------------- Wv 是 d→d 切成 H 头。第 h 个头的 v = 一个 hd×d 的矩阵把【完整的 x】压成 hd 维 —— 不是把 x 切成 H 段各拿一段。所以十个头是"十份对同一个东西的不同 角度的压缩", 这就是"自成一脉"。 ★ 但这句话有个前提没人验过: 万一十个头压的是同一个角度呢? 每个头的 hd 行张成 R^d 里一个 hd 维子空间。测两两重叠: overlap(a,b) = ‖Q_a Q_bᵀ‖²_F / hd (Q = 行空间的正交基) 随机基线 = hd/d = 0.1 (128/1280) ≈0.1 十个子空间近正交 ⇒ 十份互补的摘要, "自成一脉"成立 →1 高度重叠 ⇒ 十个头看同一片区域, 互相冗余, "自成一脉"只是名义上的 同时报每个头 Wv 块的有效秩: 若远低于 hd, 说明这个头根本没用满自己那 128 维。 B. 单头消融热图 (需要数据) ------------------------------------------------------------------------- 把第 k 层第 h 个头的输出置零, 测 VAL 涨多少。L×H 张表。 ★ 这是本架构独有的能力: 没有 Wo + 槽写一次不覆盖 ⇒ head h 永远物理占据 [h·hd,(h+1)·hd), 置零是一次精确赋值。标准 transformer 里 Wo 一乘、残差 一加, 得先"解叠加"才能动单个头, 而解得准不准本身就是个研究课题。 C. 探针 × 头 归因 (需要数据) ------------------------------------------------------------------------- 只在某类探针样本上测消融影响。挖掉哪个头会让 dup 崩掉但别的探针没事? 那个头就是归纳头。这是能直接写进结论的可解释性结果。 用法: python head_test.py --model delta_d1280_h10_l10.py \\ --ckpt /workspace/data/d1280_h10_l10_out/latest.pt \\ --batches 8 --probe-batches 6 SMOKE=1 python head_test.py --model delta_d1280_h10_l10.py --batches 2 ================================================================================ """ import argparse, importlib.util, itertools import numpy as np import torch import torch.nn.functional as F def load_module(path): spec = importlib.util.spec_from_file_location("deltamod", path) m = importlib.util.module_from_spec(spec) spec.loader.exec_module(m) return m # ═══════════════════ A. 头子空间独立性 ═══════════════════ @torch.no_grad() def head_subspaces(model): L, H = model.L, model.H hd = model.d // H rows = [] for i in range(L): W = model.atts[i].Wv.weight.data.float() # (d, d), 行=输出维 Q, erank = [], [] for h in range(H): B = W[h * hd:(h + 1) * hd, :] # (hd, d) 该头的压缩矩阵 U, S, Vh = torch.linalg.svd(B, full_matrices=False) Q.append(Vh) # 行空间正交基 (hd, d) p = S**2 / (S**2).sum() erank.append(float(torch.exp(-(p * (p + 1e-12).log()).sum()))) ov = [] for a, b in itertools.combinations(range(H), 2): ov.append(((Q[a] @ Q[b].T).pow(2).sum() / hd).item()) rows.append((np.mean(ov), np.max(ov), np.mean(erank), np.min(erank))) return rows, hd / model.d # ═══════════════════ 消融评估 ═══════════════════ @torch.no_grad() def ablate_loss(m, model, batches, patches, probe=None): """patches=[(层,头,None)] 置零。probe=None 时返回全体平均 loss; 否则返回该探针位置上的平均 loss。batches 是预取好的 (x,y) 列表。""" tot, n = 0.0, 0 with m.amp_ctx(): for x, y in batches: lg = model.intervene(x, patches) ls = F.cross_entropy(lg.reshape(-1, lg.size(-1)), y.reshape(-1), reduction="none").view_as(y) msk = (x == probe) if probe is not None else (x < 50257) if msk.any(): tot += ls[msk].sum().item() n += int(msk.sum()) return tot / max(n, 1) def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", default="delta_d1280_h10_l10.py") ap.add_argument("--ckpt", default=None) ap.add_argument("--batches", type=int, default=8) ap.add_argument("--probe-batches", type=int, default=6) ap.add_argument("--skip-probe", action="store_true") a = ap.parse_args() m = load_module(a.model) model = m.DeltaLM().to(m.device) tag = "随机初始化 (仅验证通路, 结论无意义)" if a.ckpt: ck = torch.load(a.ckpt, map_location=m.device, weights_only=False) model.load_state_dict(ck["model"]) tag = f"{a.ckpt} step={ck.get('step','?')}" model.eval() L, H = model.L, model.H print("=" * 100) print(f" 头独立性 & 单头消融归因 {tag}") print("=" * 100) # ── A ── rows, base_ov = head_subspaces(model) print(f"\n 【A】Wv 头子空间独立性 (纯权重) 随机基线 = hd/d = {base_ov:.3f}") print(" " + "-" * 96) print(f" {'层':<6}{'平均重叠':<12}{'最大重叠':<12}{'平均有效秩':<14}{'最低有效秩':<14}{'判读'}") for i, (mo, xo, me, mi) in enumerate(rows): v = ("近正交,互补" if mo < base_ov * 1.5 else "轻度冗余" if mo < base_ov * 3 else "高度冗余") print(f" L{i+1:<5}{mo:<12.4f}{xo:<12.4f}{me:<14.1f}{mi:<14.1f}{v}") print(f"\n ★ 重叠≈{base_ov:.2f} ⇒ 十个头读 x 的不同方向, '自成一脉'成立") print(f" ★ 有效秩满值={model.d//H}。远低于它 = 该头没用满自己那 {model.d//H} 维") # ── 数据 ── if m.SMOKE: vl = m.SyntheticLoader(m.MICRO_BATCH, m.SEQ_LEN, 2) pl = m.SyntheticLoader(m.MICRO_BATCH, m.SEQ_LEN, 3) else: vl = m.ShardLoader(m.DATA_DIR, "val", m.MICRO_BATCH, m.SEQ_LEN, insert_p=0.0) pl = m.ShardLoader(m.DATA_DIR, "val", m.MICRO_BATCH, m.SEQ_LEN, insert_p=0.05, seed=m.SEED + 999) vl.reset() vb = [vl.next_batch(m.device) for _ in range(a.batches)] # ── B ── base = ablate_loss(m, model, vb, []) print(f"\n 【B】单头消融热图 基线 loss = {base:.4f} ({a.batches} batch)") print(" " + "-" * 96) print(" " + "层\\头".ljust(8) + "".join(f"h{h+1:<7}" for h in range(H))) grid = np.zeros((L, H)) for k in range(1, L + 1): for h in range(H): grid[k-1, h] = ablate_loss(m, model, vb, [(k, h, None)]) - base print(f" L{k:<7}" + "".join(f"{grid[k-1,h]:<8.3f}" for h in range(H))) flat = [(grid[i, j], i + 1, j + 1) for i in range(L) for j in range(H)] flat.sort(reverse=True) print(f"\n 最关键的 5 个头: " + " | ".join(f"L{k}h{h} {v:+.3f}" for v, k, h in flat[:5])) print(f" 最没用的 5 个头: " + " | ".join(f"L{k}h{h} {v:+.3f}" for v, k, h in flat[-5:])) print(f" 全体 {L*H} 个头: 中位 {np.median(grid):+.3f} " f"占比>0.01 的 {int((grid>0.01).sum())} 个") print(" ★ 大量头 ΔVAL≈0 = 冗余, 可考虑减头或减层; 少数头独大 = 功能高度集中") # ── C ── if not a.skip_probe: pl.reset() pb = [pl.next_batch(m.device) for _ in range(a.probe_batches)] print(f"\n 【C】探针 × 头 归因 ({a.probe_batches} batch, 插入率 5%)") print(" " + "-" * 96) print(f" {'探针':<8}{'基线':<10}{'最关键的3个头 (ΔLoss)'}") for nm in m.PROBES: tok = m._PROBE_TOK[nm] b0 = ablate_loss(m, model, pb, [], probe=tok) if not np.isfinite(b0) or b0 == 0: print(f" {nm:<8}{'无样本':<10}") continue sc = [] for k in range(1, L + 1): for h in range(H): sc.append((ablate_loss(m, model, pb, [(k, h, None)], probe=tok) - b0, k, h + 1)) sc.sort(reverse=True) print(f" {nm:<8}{b0:<10.3f}" + " | ".join(f"L{k}h{h} {v:+.3f}" for v, k, h in sc[:3])) print("\n ★ 若某个头对 dup 影响极大而对其他探针几乎无影响 ⇒ 它就是归纳头。") print(" 这是本架构独有、标准 transformer 复现不了的定位结果。") print(" ★ 注意: 消融是'挖掉看塌不塌', 只证明必要性, 不证明充分性。") print("=" * 100) if __name__ == "__main__": main()