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()