File size: 18,334 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
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
#!/usr/bin/env python3
"""
================================================================================
additive_test.py — 检验槽位栈的 gate 读出是否超出"加性残差可表达类"
================================================================================

一. 要证的命题
--------------
加性残差流里 consumer 读到的是  W(Σ_j α_j s_j) = Σ_j α_j·W s_j
展开成分块矩阵就是   [α_0 W, α_1 W, …, α_{k-1} W]
    → 每一行都是同一个 u 按同一组 α 复制 k 份
    → 全部 dff 个输出坐标共用一组混合比 α

槽位栈是            [W^(0), W^(1), …, W^(k-1)]  各块彼此无关
    → 每个输出坐标可以有自己的混合比

★ 所以"加性可表达"的精确含义 = W_g 沿块维度是秩1的。
  这给出一个可判定的检验, 而不是靠聚合占比去猜。

★ 为何"聚合贡献占比偏离范数比例"证明不了:
  1) 标准残差里 W 可以挑子空间, 把某个源投到零 —— 占比与范数比例无关是
     加性残差本来就能做到的事;
  2) 聚合占比是在全部输出坐标上平均的, 恰好把"逐坐标混合比不同"这个真正
     的差别平均掉。两个模型可以聚合占比完全相同, 一个逐坐标高度分化,
     一个完全均一。

二. 三个测试
------------
  T1 逐坐标混合比离散度  纯权重分析。对每个输出坐标算它对各槽的行范数占比
                         p_m ∈ 单纯形; 加性预测所有 p_m 相同。测散度。
                         (这是"聚合占比"的正确版本: 看离散度不看均值)
  T2 加性投影            沿块维度 SVD, 秩r 截断后量 VAL。r=1 即"投影到加性
                         可表达类"。★必须配同 Frobenius 幅度的随机扰动对照,
                         否则分不清"结构被破坏"与"权重被改动太多"。
  T3 块置换              各块读的槽做错排。加性下各块只差标量, 置换近乎无损。

三. 判读
--------
  T2 是决定性的: 若 VAL(秩1) ≫ VAL(同幅度随机扰动), 则被破坏的是结构本身,
  加性残差确实表达不了当前的读出方式。若两者相当, 说明这些层的 gate 实际上
  停留在加性可表达的范围内 —— 那是个真实的负结果, 该认。

  T1/T3 是支持性证据: 它们能证明"分化了", 但不能单独证明"加性做不到"。

四. 用法
--------
    python additive_test.py --ckpt /workspace/data/d1280_h10_l10_out/latest.pt \
                            --model delta_d1280_h10_l10.py --batches 20
    SMOKE=1 python additive_test.py --model delta_d1280_h10_l10.py --batches 2
        (无 --ckpt 时用随机初始化, 只验证脚本通路; 随机初始化下各块本就正交,
         结论无意义)
================================================================================
"""
import argparse, importlib.util, math, os, sys
import numpy as np
import torch


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


# ═══════════════════════ T1: 逐坐标混合比离散度 ═══════════════════════
def coordwise_mixing(W, k, d):
    """W: (dff, k*d)。返回 (每坐标混合比矩阵 (dff,k), 平均占比 (k,), 离散度)。

    对输出坐标 m, 它对槽 j 的读出强度 = 该行在块 j 上的 L2 范数。
    归一化成单纯形上的点 p_m。加性残差预测: 所有 p_m 完全相同。

    离散度用平均总变差距离 (到全局均值) 表示, 单位 pp:
        disp = mean_m  (1/2)·Σ_j |p_m[j] - p̄[j]| × 100
    0pp = 所有坐标混合比一致 (加性形态); 越大 = 逐坐标越分化。
    """
    B = W.reshape(W.shape[0], k, d)          # (dff, k, d)  [:, j, :] = 块 j
    s = B.norm(dim=-1)                       # (dff, k)
    p = s / s.sum(-1, keepdim=True).clamp_min(1e-12)
    pbar = p.mean(0)
    disp = (0.5 * (p - pbar).abs().sum(-1)).mean().item() * 100
    return p, pbar, disp


# ═══════════════════════ T2: 加性投影 (块维 SVD) ═══════════════════════
def block_spectrum(W, k, d):
    """把 k 个块各自拉直成向量, 求 k×k Gram 的特征谱。

    加性可表达 ⟺ 这些向量共线 ⟺ 谱只有一个非零值。
    返回 (特征值降序, 秩1能量占比)。
    随机初始化基线: 各块独立随机 → 谱近似平坦 → 秩1占比 ≈ 1/k。
    """
    V = W.reshape(W.shape[0], k, d).permute(1, 0, 2).reshape(k, -1).double()
    G = V @ V.T
    ev = torch.linalg.eigvalsh(G).flip(0).clamp_min(0)
    return ev, (ev[0] / ev.sum().clamp_min(1e-30)).item()


def matrix_rank_matched(W, eps):
    """沿【普通矩阵秩】截断, 挑 rank 使 ‖ΔW‖_F 最接近 eps。

    ★ 这是比"同幅度随机扰动"严得多的对照。各向同性噪声在 dff×(k·d) 的大空间
      里落到模型实际使用的子空间上的分量极小 —— 它匹配了幅度, 却没匹配"命中
      率", 于是任何结构化删除都天然更疼, 与块结构无关。
      本对照同样是结构化低秩删除、同样的 Frobenius 幅度, 唯一区别是沿哪个轴,
      从而把"块维度"这一个因素单独隔离出来。
    """
    U, S, Vh = torch.linalg.svd(W.double(), full_matrices=False)
    tail = torch.flip(torch.cumsum(torch.flip(S**2, [0]), 0), [0]).sqrt()  # ‖Δ‖(rank r)
    r = int((tail - eps).abs().argmin().item())
    r = max(1, min(r, len(S) - 1))
    Wr = (U[:, :r] * S[:r]) @ Vh[:r]
    return Wr.to(W.dtype), r, tail[r].item()


def random_subspace_keep(W, frac_keep, gen):
    """对照D: 只保留输入空间里一个【随机】子空间, 丢掉其余。

    为何需要它: 对照B(普通矩阵秩截断)砍掉的是奇异值最小的方向 —— 同等能量下
    最良性的一种删除, 只给出损伤下界。对照D 同样是结构化删除、同样的能量,
    但方向随机挑, 不偏袒"最没用的"。它才是"一般性结构化损伤"的公允基线。

    W P, P = G(GᵀG)⁻¹Gᵀ 投影到随机 q 维子空间, q = frac_keep · (k·d)。
    各向同性下期望保留能量 ≈ q/(k·d), 故 frac_keep 直接取"块维秩1 保留的能量比"。
    """
    kd = W.shape[1]
    q = max(1, min(kd - 1, int(round(frac_keep * kd))))
    G = torch.randn(kd, q, generator=gen).to(W.device).float()
    A = W.float() @ G                              # (dff, q)
    M = G.T @ G                                    # (q, q)
    B = torch.linalg.solve(M + 1e-6 * torch.eye(q, device=W.device), A.T).T
    return (B @ G.T).to(W.dtype)


def rotated_block_rank1(W, k, d, gen):
    """对照C: 每块输入空间先乘一个随机正交阵 R_j, 做块维秩1, 再转回来。

    结果形如  W^(j) = α_j · U R_jᵀ
      · 仍是"秩1 + 全局标量 α"  —— 逐坐标幅度混合比与加性残差同样受限
      · 但每个槽被读的方向不同  —— 加性残差表达不了这一点

    于是它把两个自由度拆开了:
      对照C 也很疼  ⇒ 关键自由度是"逐坐标幅度混合比"
      对照C 良性    ⇒ 关键自由度是"从各源读不同方向", 与幅度分配无关

    ⚠ 实测判定本对照【无效】, 保留仅供复现: 随机旋转等于"先把各槽内容打乱
      再做加性合并", 是"加性残差 + 输入被搅烂", 严格劣于加性本身, 拿它对照
      证明不了任何事。更根本的问题是"加性可表达" ⟺ "块维秩1" 是同一个集合
      (V 秩1 ⟹ W^(j)=α_j U 强制成立), 无法固定秩而只改加性与否。
      公允对照请用 random_subspace_keep (对照D)。
    """
    dff = W.shape[0]
    B = W.reshape(dff, k, d)
    Rs, Bt = [], torch.empty_like(B)
    for j in range(k):
        Q, _ = torch.linalg.qr(torch.randn(d, d, generator=gen).to(W.device))
        Rs.append(Q.to(W.dtype))
        Bt[:, j, :] = B[:, j, :] @ Rs[j]
    B1 = rank_truncate(Bt.reshape(dff, k * d), k, d, 1).reshape(dff, k, d)
    out = torch.empty_like(B)
    for j in range(k):
        out[:, j, :] = B1[:, j, :] @ Rs[j].T
    return out.reshape(dff, k * d)


def rank_truncate(W, k, d, r):
    """沿块维度做秩 r 截断。r=1 即投影到加性残差可表达类。"""
    dff = W.shape[0]
    V = W.reshape(dff, k, d).permute(1, 0, 2).reshape(k, -1).double()
    U, S, Vh = torch.linalg.svd(V, full_matrices=False)
    Vr = (U[:, :r] * S[:r]) @ Vh[:r]
    return Vr.reshape(k, dff, d).permute(1, 0, 2).reshape(dff, k * d).to(W.dtype)


# ═══════════════════════ 评估 ═══════════════════════
@torch.no_grad()
def val_loss(m, model, loader, nb):
    model.eval()
    loader.reset()
    t = 0.0
    with m.amp_ctx():
        for _ in range(nb):
            x, y = loader.next_batch(m.device)
            t += model(x, y)[1].item()
    return t / nb


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=20)
    ap.add_argument("--ranks", default="1,2,3")
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--per-layer", action="store_true",
                    help="每次只对一层做块维秩1截断, 跳出饱和区并做逐层归因")
    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()

    if m.SMOKE:
        loader = m.SyntheticLoader(m.MICRO_BATCH, m.SEQ_LEN, 2)
    else:
        loader = m.ShardLoader(m.DATA_DIR, "val", m.MICRO_BATCH, m.SEQ_LEN,
                               insert_p=0.0)

    L, d, dff = model.L, model.d, m.DFF_
    print("=" * 100)
    print(f"  加性可表达性检验   {tag}")
    print(f"  d={d} dff={dff} L={L}   VAL 用 {a.batches} 个 batch")
    print("=" * 100)

    base = val_loss(m, model, loader, a.batches)
    print(f"\n  基线 VAL = {base:.4f}\n")

    # ── T1 + 谱 ──
    print("  【T1】逐坐标混合比离散度 & 块维谱")
    print("  " + "-" * 96)
    print(f"  {'层':<5}{'k':<4}{'离散度(pp)':<13}{'秩1能量占比':<14}"
          f"{'随机基线1/k':<13}{'谱(前4, 归一化)'}")
    for i in range(1, L):                    # ffn_0 只有 1 块, 平凡加性
        W = model.ffns[i].Wg.weight.data.float()
        k = i + 1
        _, pbar, disp = coordwise_mixing(W, k, d)
        ev, r1 = block_spectrum(W, k, d)
        evn = (ev / ev.sum()).tolist()[:4]
        print(f"  f{i+1:<4}{k:<4}{disp:<13.1f}{r1:<14.4f}{1.0/k:<13.4f}"
              + " ".join(f"{v:.3f}" for v in evn))
    print("\n  离散度 0pp = 所有输出坐标共用一组混合比 = 加性形态")
    print("  秩1能量占比 →1 = 各块共线 = 加性可表达; →1/k = 各块正交")
    print("  ★ 随机初始化时秩1占比本就≈1/k, 故该列低不能单独当作学到了东西")

    # ── T2 ──
    ranks = [int(r) for r in a.ranks.split(",")]
    print(f"\n  【T2】加性投影 (沿块维秩r截断) + 同幅度随机扰动对照")
    print("  " + "-" * 96)
    orig = {i: model.ffns[i].Wg.weight.data.clone() for i in range(L)}
    g = torch.Generator(device="cpu").manual_seed(a.seed)

    for r in ranks:
        # 秩 r 截断
        dfro = 0.0
        for i in range(1, L):
            W = orig[i].float()
            Wr = rank_truncate(W, i + 1, d, r)
            dfro += (Wr - W).pow(2).sum().item()
            model.ffns[i].Wg.weight.data.copy_(Wr.to(orig[i].dtype))
        lo = val_loss(m, model, loader, a.batches)
        for i in range(L):
            model.ffns[i].Wg.weight.data.copy_(orig[i])

        # 对照: 同 Frobenius 幅度的随机扰动, 逐层按各自幅度分配
        for i in range(1, L):
            W = orig[i].float()
            Wr = rank_truncate(W, i + 1, d, r)
            eps = (Wr - W).norm()
            n = torch.randn(W.shape, generator=g).to(W.device)
            n = n / n.norm() * eps
            model.ffns[i].Wg.weight.data.copy_((W + n).to(orig[i].dtype))
        lc = val_loss(m, model, loader, a.batches)
        for i in range(L):
            model.ffns[i].Wg.weight.data.copy_(orig[i])

        # 严对照: 同幅度的【普通矩阵秩】截断 (结构化删除, 只是换了个轴)
        rs = []
        for i in range(1, L):
            W = orig[i].float()
            eps = (rank_truncate(W, i + 1, d, r) - W).norm()
            Wm, rr, _ = matrix_rank_matched(W, eps)
            rs.append(rr)
            model.ffns[i].Wg.weight.data.copy_(Wm.to(orig[i].dtype))
        lm = val_loss(m, model, loader, a.batches)
        for i in range(L):
            model.ffns[i].Wg.weight.data.copy_(orig[i])

        print(f"  块维秩{r}: VAL {lo:.4f}{lo-base:+.4f})  ‖ΔW‖_F={math.sqrt(dfro):.1f}")
        print(f"      对照A 同幅度随机扰动     VAL {lc:.4f}{lc-base:+.4f})   [弱对照]")
        print(f"      对照B 同幅度普通矩阵秩截断 VAL {lm:.4f}{lm-base:+.4f})   "
              f"[严对照, matrix rank≈{int(np.mean(rs))}]")
        if lo > lm + 0.05:
            print(f"      ⇒ 块维删除比同幅度普通秩删除更疼 (Δ差 {lo-lm:+.3f}), "
                  f"块结构被单独隔离出来了")
        else:
            print(f"      ⇒ 两者相当, 损伤不能归因于块结构本身")
    print("\n  ★ 判据以【对照B】为准: 块维秩1 显著疼于同幅度普通矩阵秩截断")
    print("    ⇒ 损伤来自块结构本身, 加性残差表达不了当前读出方式。")
    print("  ★ 注意各截断若都落在'已损坏'区间(VAL≈ln(vocab)), 秩阶梯之间的排序")
    print("    是噪声, 只有'秩1 vs 对照'这一比较有效。")
    print("  ★ 本测试证明的是: 该已训练模型重度使用了加性表达不了的成分。它")
    print("    没有证明'原生训练一个加性模型也达不到同样 loss' —— 后者需要重训")
    print("    tied-block 基线, 且有严重参数量混淆。这是不重训能拿到的最强证据。")

    # ── T2b: 逐层归因 ──
    if a.per_layer:
        print(f"\n  【T2b】逐层块维秩1截断 (一次只动一层, 跳出饱和区)")
        print("  " + "-" * 96)
        print(f"  {'层':<6}{'k':<4}{'块维秩1(加性)':<16}{'对照B 最良性删除':<20}"
              f"{'对照D 随机子空间':<20}{'‖Δ‖块':<9}{'‖Δ‖D':<9}{'保留能量'}")
        rows = []
        for i in range(1, L):
            W = orig[i].float()
            Wr = rank_truncate(W, i + 1, d, 1)
            eps = (Wr - W).norm()
            model.ffns[i].Wg.weight.data.copy_(Wr.to(orig[i].dtype))
            lb = val_loss(m, model, loader, a.batches)
            Wm, rr, _ = matrix_rank_matched(W, eps)
            model.ffns[i].Wg.weight.data.copy_(Wm.to(orig[i].dtype))
            lc2 = val_loss(m, model, loader, a.batches)
            model.ffns[i].Wg.weight.data.copy_(orig[i])
            fk = 1.0 - (eps / W.norm()) ** 2          # 块维秩1 保留的能量比
            Wd_ = random_subspace_keep(W, float(fk), g)
            epsd = (Wd_ - W).norm().item()
            model.ffns[i].Wg.weight.data.copy_(Wd_.to(orig[i].dtype))
            lc4 = val_loss(m, model, loader, a.batches)
            model.ffns[i].Wg.weight.data.copy_(orig[i])
            rows.append((i + 1, lb - base, lc2 - base, lc4 - base))
            print(f"  f{i+1:<5}{i+1:<4}{lb-base:<16.4f}{lc2-base:<20.4f}"
                  f"{lc4-base:<20.4f}{eps.item():<9.1f}{epsd:<9.1f}{fk*100:.0f}%")
        bb = np.array([r[1] for r in rows]); dd = np.array([r[3] for r in rows])
        print(f"\n  合计 ΔVAL: 块维秩1(加性) {bb.sum():.2f} | 对照B(最良性) "
              f"{np.array([r[2] for r in rows]).sum():.2f} | 对照D(随机子空间) {dd.sum():.2f}")
        print("\n  ★ 判据看【对照D】: 同样丢掉同等能量, 但丢的方向随机挑。")
        if bb.sum() > 1.5 * dd.sum():
            print("    块维秩1 明显更疼 ⇒ 损伤不是'丢了这么多能量'的一般后果,")
            print("    而是特定于'塌成加性形态'这个结构。结论成立。")
        elif dd.sum() > 1.5 * bb.sum():
            print("    随机子空间反而更疼 ⇒ 加性形态其实保住了不少功能,")
            print("    之前对照B 得出的差距是假象。结论需推翻。")
        else:
            print("    两者相当 ⇒ 损伤主要来自'丢了这么多能量', 与是否加性无关。")
            print("    这是真实的负结果, 该认。")
        print("  ★ 对照B 砍的是奇异值最小的方向, 是同能量下最良性的删除, 只给")
        print("    损伤下界; 与它的差距被系统性放大, 倍数别当真。")

    # ── T3 ──
    print(f"\n  【T3】块置换 (错排, 各块改读别的槽)")
    print("  " + "-" * 96)
    rng = np.random.default_rng(a.seed)
    for trial in range(3):
        for i in range(1, L):
            k = i + 1
            while True:                       # 错排: 无不动点
                perm = rng.permutation(k)
                if not (perm == np.arange(k)).any():
                    break
            W = orig[i].reshape(dff, k, d)
            model.ffns[i].Wg.weight.data.copy_(
                W[:, torch.from_numpy(perm.copy()), :].reshape(dff, k * d))
        lp = val_loss(m, model, loader, a.batches)
        for i in range(L):
            model.ffns[i].Wg.weight.data.copy_(orig[i])
        print(f"  错排 #{trial+1}: VAL {lp:.4f}{lp-base:+.4f})")
    print("\n  ★ 加性形态下各块只差一个标量, 置换应近乎无损。VAL 大幅上升")
    print("    ⇒ 每块已专门化到自己那个槽。(支持性证据, 不能单独证明加性做不到)")
    print("=" * 100)


if __name__ == "__main__":
    main()