| |
| """ |
| ================================================================================ |
| 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): |
| |
| |
| G = np.full((L, L), np.nan) |
| for i in range(1, L): |
| for j in range(i + 1): |
| 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))) |
|
|
| |
| |
| |
| |
| 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))) |
|
|
| |
| 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 历史 ⇒ 可大幅省参") |
|
|
| |
| |
| |
| |
| |
| 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() |
|
|