File size: 9,051 Bytes
bbb9a9b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | #!/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()
|