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