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