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