| |
| """ |
| ================================================================================ |
| 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 |
|
|
|
|
| |
| @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() |
| Q, erank = [], [] |
| for h in range(H): |
| B = W[h * hd:(h + 1) * hd, :] |
| U, S, Vh = torch.linalg.svd(B, full_matrices=False) |
| Q.append(Vh) |
| 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) |
|
|
| |
| 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)] |
|
|
| |
| 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 = 冗余, 可考虑减头或减层; 少数头独大 = 功能高度集中") |
|
|
| |
| 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() |
|
|