File size: 9,524 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
#!/usr/bin/env python3
"""
================================================================================
source_test.py — 槽间因果依赖矩阵: "第 i 层真的在读第 j 层吗"
================================================================================

一. 为什么这个实验只有本架构做得了
----------------------------------
ffn_i 的 gate 是 k 个【物理分开】的矩阵块, 第 j 块专门读第 j 个槽:
    gate = [W^(0) W^(1) … W^(k-1)] · concat(槽0 … 槽k-1)
所以"第 7 层到底有没有真的在用第 3 层的输出"可以直接把 W^(3) 清零来测。

标准 transformer 里没有可清零的对象 —— 所有层的信息叠在一条残差流里,
共用一个读出矩阵, "读第 3 层的那块权重"这个东西不存在。

★ 与训练日志里 gate依赖 那一行的区别:
    日志里的百分比 = 各块贡献的范数占比 = 【相关性】
    本脚本 = 清零该块后 loss 涨多少        = 【因果】
  两者可能差很远: 某块占比 12% 但清零无损 (读的是冗余信息), 或占比 6%
  却清零就崩。只有后者能说明"这条边真的在承载信息"。

二. 输出
--------
  一张 L×L 下三角图 (行=消费者 ffn_i, 列=来源槽 j)。
  对角线右侧为空 —— ffn_i 的 gate 只读槽 0..i-1 (最新的槽 i 走 up 那一路)。
  另附 up 通路消融: 把 Wu 清零, 看"最新 v 负责给"这一路的权重。

三. 用法
--------
    python source_test.py --model delta_d1280_h10_l10.py \\
        --ckpt /workspace/data/d1280_h10_l10_out/latest.pt --batches 20
    加 --probes 同时给出每个探针各自的依赖矩阵 (慢 7 倍)
================================================================================
"""
import argparse, importlib.util
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 loss_on(m, model, batches, probe=None):
    tot, n = 0.0, 0
    with m.amp_ctx():
        for x, y in batches:
            lg = model(x)[0]
            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=20)
    ap.add_argument("--probes", 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, d = model.L, model.d

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

    print("=" * 100)
    print(f"  槽间因果依赖矩阵   {tag}")
    print("=" * 100)

    orig = {i: model.ffns[i].Wg.weight.data.clone() for i in range(L)}
    origu = {i: model.ffns[i].Wu.weight.data.clone() for i in range(L)}

    def run_matrix(probe, batches, base):
        # ★ i=0 (f1) 只有一个 gate 块, 清零它 = gate全零 = 输出全零 = 整个模型死掉。
        #   那不是一条"依赖边"而是"杀掉模型", 数值会淹没所有真实的边, 故跳过。
        G = np.full((L, L), np.nan)
        for i in range(1, L):
            for j in range(i + 1):                    # ffn_i 读槽 0..i
                W = model.ffns[i].Wg.weight.data
                W[:, j * d:(j + 1) * d] = 0
                G[i, j] = loss_on(m, model, batches, probe) - base
                model.ffns[i].Wg.weight.data.copy_(orig[i])
        return G

    base = loss_on(m, model, vb)
    print(f"\n  基线 loss = {base:.4f}   ({a.batches} batch)")
    G = run_matrix(None, vb, base)

    print(f"\n  【因果依赖】清零 ffn_i 读槽 j 的那块权重后 loss 涨幅")
    print("  (f1 跳过: 它只有一块, 清零=杀掉整个模型, 不是依赖边)")
    print("  " + "-" * 96)
    hdr = "  " + "消费者\\来源".ljust(12) + "x".ljust(8) + "".join(
        f"b{j}".ljust(8) for j in range(1, L))
    print(hdr)
    for i in range(L):
        row = "  " + f"f{i+1}".ljust(12)
        for j in range(L):
            row += ("".ljust(8) if np.isnan(G[i, j]) else f"{G[i,j]:.3f}".ljust(8))
        print(row)

    print(f"\n  列合计 (该槽被下游总共依赖多少):")
    cs = np.nansum(G, axis=0)
    print("    " + "  ".join(f"{('x' if j==0 else 'b'+str(j))}:{cs[j]:.2f}"
                             for j in range(L)))
    print(f"  行合计 (该 ffn 对历史的总依赖):")
    rs = np.nansum(G, axis=1)
    print("    " + "  ".join(f"f{i+1}:{rs[i]:.2f}" for i in range(L)))

    # ── gate 是否需要随输入变化 ──
    # ★ 单块清零"几乎无损"推不出"这些块没用": 可能每块单独可省, 合起来才重要。
    #   直接测: 把 Wg 各行设成同一向量 ⇒ gate 变成常数向量(不随 token 变),
    #   幅度保留, 但"历史负责选"这个机制被整体抹掉。
    print(f"\n  【gate 是否随输入变化】把 Wg 各行设为同一向量 (抹掉选择机制)")
    print("  " + "-" * 96)
    gc = []
    for i in range(L):
        W = model.ffns[i].Wg.weight.data
        W.copy_(W[:1].expand_as(W).contiguous())
        gc.append(loss_on(m, model, vb) - base)
        model.ffns[i].Wg.weight.data.copy_(orig[i])
    print("    " + "  ".join(f"f{i+1}:{v:.3f}" for i, v in enumerate(gc)))

    # ── 只保留 x 块, 抹掉全部 att 历史 ──
    print(f"\n  【只保留 x 块】清零 ffn_i 读槽 b1..b(i-1) 的全部权重")
    print("  " + "-" * 96)
    xo = []
    for i in range(1, L):
        W = model.ffns[i].Wg.weight.data
        W[:, d:(i + 1) * d] = 0
        xo.append(loss_on(m, model, vb) - base)
        model.ffns[i].Wg.weight.data.copy_(orig[i])
    print("    " + "  ".join(f"f{i+2}:{v:.3f}" for i, v in enumerate(xo)))
    print("  ★ 三个数一起读:")
    print("    单块清零≈0 且 常数化≈0 且 只保留x≈0  ⇒ 该层 gate 确实是摆设")
    print("    单块清零≈0 但 常数化很大            ⇒ 冗余, 各块互为备份, 不能删")
    print("    只保留x≈0 但 常数化很大             ⇒ gate 只需要 token 身份,")
    print("                                          不需要 att 历史 ⇒ 可大幅省参")

    # up 通路
    # ★ 不能用"清零 Wu": up 出来要过 N_ 硬归一化, N_(0)=0 ⇒ 整个 ffn 输出全零
    #   ⇒ 下游全部塌成同一个退化态, 不管动哪一层结果都一模一样, 毫无信息。
    #   正确做法: 把 Wu 的所有行设成同一个向量 ⇒ up 各维取值相同 ⇒ 过 N_ 后
    #   变成常数向量。内容信息被抹掉, 幅度 ‖up‖≡√dff 保持不变。
    print(f"\n  【up 通路】把 Wu 各行设为同一向量 (抹掉'给'的内容, 保幅度)")
    print("  " + "-" * 96)
    ups = []
    for i in range(L):
        W = model.ffns[i].Wu.weight.data
        W.copy_(W[:1].expand_as(W).contiguous())
        ups.append(loss_on(m, model, vb) - base)
        model.ffns[i].Wu.weight.data.copy_(origu[i])
    print("    " + "  ".join(f"f{i+1}:{v:.3f}" for i, v in enumerate(ups)))
    print("  ★ 这一路是'最新 att 输出负责给内容'。数值远大于 gate 单块 = 该层的")
    print("    内容供给比任何单条历史读边都重要 (符合设计意图: 历史选, 最新给)。")

    if a.probes:
        pl.reset()
        pb = [pl.next_batch(m.device) for _ in range(max(4, a.batches // 2))]
        print(f"\n  【分探针】各探针最依赖的 3 条边 (消费者←来源)")
        print("  " + "-" * 96)
        for nm in m.PROBES:
            tok = m._PROBE_TOK[nm]
            b0 = loss_on(m, model, pb, tok)
            if b0 <= 0:
                print(f"  {nm:<8} 无样本或已完全解决")
                continue
            Gp = run_matrix(tok, pb, b0)
            fl = sorted(((Gp[i, j], i + 1, j) for i in range(1, L)
                         for j in range(i + 1)), reverse=True)
            print(f"  {nm:<8}基线{b0:<8.3f}" + " | ".join(
                f"f{i}{'x' if j==0 else 'b'+str(j)} {v:+.3f}" for v, i, j in fl[:3]))
        print("\n  ★ 这给出每个功能各自的信息流路径。若 dup 高度依赖 f8←b7 而")
        print("    b128 高度依赖 f3←b2, 说明两条电路在深度上是分离的。")

    print("\n  ★ 与训练日志 gate依赖 百分比对照: 那个是范数占比(相关性), 本表是")
    print("    清零后的 loss 涨幅(因果)。占比高但清零无损 = 读的是冗余信息。")
    print("=" * 100)


if __name__ == "__main__":
    main()