稀疏权重分解与电路提取:从神经网络中还原真实计算路径

📅 2026/8/27 4:25:47
稀疏权重分解与电路提取:从神经网络中还原真实计算路径
如果你手里有一个已经训练好的神经网络想弄清楚它到底“记住”了哪些计算路径你会怎么做过去最常见的答案是注意力热力图、梯度归因、激活值可视化。但这些方法只能告诉你“哪里重要”没法告诉你“信号是怎么流过去的”。在自动化机器学习、反欺诈规则提炼、芯片辅助设计这类场景里我们需要的是结构级解释把网络内部的真实计算子图抽出来看信息沿着哪些边流动。这就是 circuit extraction电路提取。真正上手做电路提取时你会发现第一个拦路虎不是图算法而是权重矩阵。一个普通的多层感知机权重连接数动辄几十万里面大量绝对值接近零的连接会制造出密集的伪路径。如果不先对权重做分解和稀疏化提取出来的“电路”和原网络没什么区别甚至比原网络更难分析。所以这个领域里出现了一个很务实的组合稀疏权重分解 电路提取。先通过分解把权重矩阵压缩成稀疏的显著结构再在稀疏结构上运行图搜索速度和可解释性同时上来。这篇文章会从原理讲到最小可运行示例最终带你跑通这样一条链路加载或训练一个简单 MLP - 用 SVD 做低秩近似 - 阈值稀疏化 - 用 networkx 构建可遍历的计算子图 - 验证提取后的子图行为和原模型是否一致。1. 这篇文章真正要解决的问题先说一个反直觉的判断电路提取的瓶颈通常不是下游的图分析而是上游的权重预处理。如果你想从神经网络中抽取“电路”首先要定义什么是“边”。在 MLP 里边就是一个权重连接在卷积网络里边可能对应某个卷积核或者某个通道。问题在于训练良好的网络权重往往不是自然稀疏的几乎所有连接都有非零权重只是绝对值大小不同。这意味着直接按权重大小剪枝容易保留碎片化路径导致提取出的子图不连通直接拿全量权重建图边数爆炸且显著性信息被大量弱连接淹没只做稀疏化但不做低秩分解会把真正有趣的“协同激活结构”破坏掉。稀疏权重分解解决的是这一层问题。它的核心思路是把权重矩阵 W 分解成若干结构性成分只保留能量占比最高的那部分成分再对结果做稀疏化。这样得到的权重既保留了主要信息通路又足够稀疏适合作为电路提取的输入。文章适合以下读者正在做模型可解释性研究的工程师想把“行为解释”升级为“结构解释”做模型压缩和加速的开发者想从权重结构中找到可裁剪的高价值模块需要在生产环境对模型做审计、规则提炼或辅助决策解释的团队。读完你会得到两条东西一是理解稀疏权重分解到底怎么和电路提取衔接二是一套可以直接扩展到你自己的网络上的 Python 示例代码。2. 稀疏权重分解与电路提取的基础概念2.1 权重分解不是压缩而是“结构显影”权重分解是一族方法的统称常见的包括 SVD奇异值分解、NMF非负矩阵分解、低秩近似等。以 SVD 为例对于形状为(out_features, in_features)的权重矩阵 W它可以写为W ≈ U_k · S_k · V_k^T其中U_k是out_features x k的正交矩阵S_k是k个奇异值组成的对角向量V_k^T是k x in_features的矩阵。k 通常远小于 out_features 和 in_features这就是低秩近似。但注意SVD 做出来的近似结果仍然是稠密矩阵。它只是把矩阵的“信息”压缩到几个主成分上并没有让矩阵里的连接变成零。所以 SVD 通常要跟稀疏化配合要么对分解结果做阈值裁剪要么在分解后的子空间里再施加结构约束。这就是标题里 “Sparse Weight Decomposition” 的含义——不是简单做稀疏化而是“分解后再稀疏化”。2.2 稀疏化结构稀疏优于数值稀疏说到稀疏化很多人的第一反应是 magnitude pruning也就是把绝对值小于阈值的权重直接置零。这种方式确实能得到稀疏矩阵但问题在于置零的位置是零散的后续在图结构上做路径搜索时这些零散断点会产生大量不连通结构。更推荐的是“能量导向的结构化稀疏”。先计算每个成分的贡献比如奇异值的平方和决定保留哪些成分然后再把保留成分内部的弱连接裁掉。这种做法的好处是被保留的边在能量意义上是重要的而不是单纯在数值上够大。2.3 电路提取从“权重集合”到“计算子图”电路提取指的是从神经网络中还原出信息流动路径。常见做法是把网络看成一张有向图节点是神经元卷积场景里也可以是通道边是权重连接边的方向代表信号传播方向。在这样一张图上我们可以回答很多有意思的问题从输入到某一个输出神经元经过的 Top-K 路径是什么哪些中间神经元同时被多条重要路径共享某个任务对应的“核心子网络”长什么样如果没有稀疏化这张图会稠密得让人无法直视。而稀疏权重分解介入之后图规模可能缩小到原来的十分之一甚至更小这才让路径搜索具备工程可行性。方法输出适合回答的问题主要缺点注意力权重可视化热力图哪些位置重要无法说明路径梯度归因输入重要性分数输入特征对输出的影响局部梯度噪声大直接全量建图完整稠密图理论上的全量路径边爆炸没法看稀疏权重分解 电路提取稀疏计算子图信息从哪里流过需要选择分解和稀疏化参数2.4 核心洞察分解让“弱连接噪音”退场真正让稀疏权重分解对电路提取有效的原因不是它减少了要看的权重数量而是它把“协同激活”从“单点刺激”里区别出来了。举个例子W 的第 i 行代表第 i 个输出神经元对全部输入神经元的重要度。SVD 会把这一行重新表达成多个成分第一个成分可能对应一组一起激活的输入特征第二个成分可能对应另一组。一个具体连接是否保留取决于它所在的成分贡献有多大。这比单独看某个权重的绝对值要合理得多。3. 核心流程拆解整个流程可以拆成五步下面逐个说明。3.1 第一步获取一个可用的神经网络电路提取面向的是已经训练好的模型。你可以用自己的模型也可以用公开数据集训练一个小模型。最小验证场景是 MLP因为它结构简单稀疏分解后可以直接映射到有向图便于观察和理解。如果你想验证方法是否在更大模型上有效建议从卷积网络或 Transformer 的某一层开始而不是一次性分析整个网络。3.2 第二步逐层做权重分解把每一层的二维权重矩阵取出来做 SVD。这里有两个选择逐层分解每层独立设置阈值全局分解把所有权重堆叠成一个大矩阵后分解。推荐从逐层分解起步。原因很简单不同层的奇异值分布差异很大输入层通常是高秩的因为要保留原始特征的多样性输出层往往可以压得很低因为类别信息集中。用同一组参数处理所有层容易过压缩某些层、欠压缩另一些层。3.3 第三步能量阈值截断 稀疏化SVD 之后每个奇异值平方占总能量平方和的比例就是该成分的信息占比。你可以设置一个能量保留率比如 0.9表示“保留 90% 的权重能量”。然后对低秩近似矩阵做硬阈值稀疏化把绝对值排在后面的边裁掉。3.4 第四步映射回原始网络构造计算图得到每层的稀疏权重后需要把权重下标映射到原有网络节点上。比如原网络输入维度是 784隐藏层是 256那么输入层节点名可以是input_0到input_783隐藏层节点名是fc1_0到fc1_255。对每一行如果某个权重大于阈值就在对应节点之间加一条有向边。3.5 第五步验证行为保真度提取电路之后不能只看图结构漂亮不漂亮还得验证这个子图的行为是否和原模型一致。比较简单的做法是把原始权重替换成稀疏化后的权重重新跑一遍推理比对输出向量的余弦相似度。如果相似度掉得太快说明稀疏化参数太激进需要回调能量保留率或阈值。4. 环境准备与前置条件本文示例使用 Python PyTorch NetworkX都是日常使用较多的机器学习标准库。版本方面不写死建议以实际项目环境为准下面给出的是足够新的组合Python 3.9 或更高PyTorch 2.xNetworkX 2.8 或更高torchvision用于 MNIST 数据集numpy。安装命令pip install torch torchvision networkx numpy如果你的机器有 GPUPyTorch 会自动使用 GPU但本文的示例模型很小CPU 跑完全没问题。数据集使用 MNIST。首次运行会自动下载如果网络环境受限可以先手动下载到./data目录或者换成自己的随机数据做演示。随机数据在验证流程里也能跑只是看不到实际的分类效果。5. 核心代码实现下面所有代码都可以保存到一个 Python 文件中从上到下运行。我会把关键点拆开讲解。5.1 代码一训练一个简单 MLP 作为分析对象文件路径train_model.pyimport torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms class SimpleMLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(28 * 28, 256) self.fc2 nn.Linear(256, 64) self.fc3 nn.Linear(64, 10) self.act nn.ReLU() def forward(self, x): x x.view(x.size(0), -1) x self.act(self.fc1(x)) x self.act(self.fc2(x)) return self.fc3(x) def load_or_train_model(epochs3): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( ./data, trainTrue, downloadTrue, transformtransform ) train_loader DataLoader(train_dataset, batch_size256, shuffleTrue) model SimpleMLP() optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.CrossEntropyLoss() model.train() for epoch in range(epochs): total_loss 0.0 for x, y in train_loader: optimizer.zero_grad() out model(x) loss loss_fn(out, y) loss.backward() optimizer.step() total_loss loss.item() print(fepoch{epoch 1}, avg_loss{total_loss / len(train_loader):.4f}) return model if __name__ __main__: model load_or_train_model() torch.save(model.state_dict(), ./model.pt)这段代码的作用是快速得到一个可以继续操作的模型。训练 3 轮后MNIST 准确率不会太高但足够验证电路提取的流程。实际项目中你应该替换成自己的模型权重。5.2 代码二实现稀疏权重分解文件路径sparse_decompose.pyimport torch def sparse_weight_decompose(weight, energy_ratio0.9, sparsity_ratio0.7, top_kNone): 对单个权重矩阵做稀疏分解。 参数: weight: (out_features, in_features) 的二维权重矩阵 energy_ratio: 保留的奇异值能量比例 sparsity_ratio: 低秩近似后要置零的比例 top_k: 如果指定则强制保留前 top_k 个奇异值 返回: sparse_weight: 稀疏化后的权重矩阵 n_components: 实际保留的奇异值个数 singular_values: 原始奇异值向量 U, S, Vh torch.linalg.svd(weight, full_matricesFalse) total_energy (S ** 2).sum() cum_energy torch.cumsum(S ** 2, dim0) if top_k is None: n_components int((cum_energy / total_energy energy_ratio).sum()) 1 n_components min(n_components, S.numel()) else: n_components min(top_k, S.numel()) U_k U[:, :n_components] S_k S[:n_components] Vh_k Vh[:n_components, :] # 低秩近似 approx U_k * S_k.unsqueeze(0) Vh_k # 硬阈值稀疏化 threshold torch.quantile(approx.abs().flatten(), sparsity_ratio) sparse_weight torch.where( approx.abs() threshold, approx, torch.zeros_like(approx), ) return sparse_weight, n_components, S def decompose_model(model, energy_ratio0.9, sparsity_ratio0.7): 遍历模型的二维权重层返回稀疏权重列表。 注意这里不修改原始模型参数。 sparse_weights [] meta [] for name, param in model.named_parameters(): if param.dim() 2 and weight in name: sw, n_comp, S sparse_weight_decompose( param.data, energy_ratioenergy_ratio, sparsity_ratiosparsity_ratio, ) sparse_weights.append(sw) meta.append({name: name, n_components: n_comp, rank: param.shape}) return sparse_weights, meta两个参数是关键energy_ratio控制低秩近似的压缩强度数值越大保留的奇异值越多越接近原矩阵sparsity_ratio控制在低秩结果上置零的比例。0.7 表示把绝对值最小的 70% 置零。这里真正值得留意的点是阈值是在低秩近似矩阵上计算而不是在原矩阵上。原因是低秩近似本身已经去掉了大量噪声成分哪怕是一个绝对值很低的系数它在主成分空间里也可能承担着结构性作用。直接对原矩阵做 magnitude pruning 没有这层保护。5.3 代码三把稀疏权重映射成计算电路图文件路径extract_circuit.pyimport networkx as nx import torch def extract_circuit(model, sparse_weights, top_edge_ratio0.2): 从稀疏权重中构建有向计算图。 参数: model: 原始 PyTorch 模型用于获取模块名称 sparse_weights: sparse_weight_decompose 返回的稀疏权重列表 top_edge_ratio: 每条边保留的比例按权重绝对值分位数计算 G nx.DiGraph() layer_order [] for name, module in model.named_modules(): if isinstance(module, torch.nn.Linear): layer_order.append(name) # 输入节点第一个全连接层的输入维度 first_linear getattr(model, layer_order[0]) in_features first_linear.in_features input_nodes [finput_{i} for i in range(in_features)] G.add_nodes_from(input_nodes, layerinput) prev_nodes input_nodes for layer_name, sparse_weight in zip(layer_order, sparse_weights): out_features, in_features sparse_weight.shape cur_nodes [f{layer_name}_{j} for j in range(out_features)] G.add_nodes_from(cur_nodes, layerlayer_name) # 计算本层边的阈值 flat sparse_weight.abs().flatten() if flat.numel() 0: threshold 0.0 else: threshold torch.quantile(flat, 1.0 - top_edge_ratio) for j in range(out_features): for i in range(in_features): w sparse_weight[j, i].item() if abs(w) threshold: G.add_edge(prev_nodes[i], cur_nodes[j], weightw) prev_nodes cur_nodes return G这段代码的思路并不复杂但有三个工程细节值得你注意节点命名要和神经网络的真实结构一一对应否则下游路径解释无法映射回原始网络top_edge_ratio0.2表示每层只保留绝对值最高的 20% 边实际项目中可以先观察层的能量分布再决定保留比例边上的weight属性保存了原始权重值后续做路径加权分析时可以直接复用而不需要重新查矩阵。5.4 代码四行为一致性验证文件路径validate_circuit.pyimport torch import torch.nn as nn def forward_with_sparse_weights(model, x, sparse_weights): 使用稀疏权重手工执行前向传播避免修改模型参数。 sparse_weights 顺序与模型中 Linear 层的顺序一致。 h x.view(x.size(0), -1) weight_iter iter(sparse_weights) for name, module in model.named_modules(): if isinstance(module, nn.Linear): W next(weight_iter).to(x.device) b module.bias h h W.T b # 除了最后的输出层其余层都经过 ReLU if name ! fc3: h torch.relu(h) return h def evaluate_consistency(model, sparse_weights, loader, max_batches2): 返回原始模型输出与稀疏权重输出之间的平均余弦相似度。 device next(model.parameters()).device model.eval() cos_sum 0.0 batch_count 0 with torch.no_grad(): for x, _ in loader: x x.to(device) out_orig model(x) out_decomp forward_with_sparse_weights(model, x, sparse_weights) cos_sim torch.cosine_similarity(out_orig, out_decomp, dim1) cos_sum cos_sim.mean().item() batch_count 1 if batch_count max_batches: break return cos_sum / max(1, batch_count)这里的验证逻辑是原模型跑一遍稀疏权重替换后的手动前向再跑一遍比较输出向量是否还保持在同一方向。余弦相似度达到 0.95 以上视为行为基本保真落在 0.85 到 0.95 之间说明稀疏化有可见影响需要检查参数是否过激低于 0.85通常就是能量保留率或稀疏比例太激进。6. 运行结果与效果验证把上面的代码串起来在main.py中执行import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms from train_model import SimpleMLP, load_or_train_model from sparse_decompose import decompose_model from extract_circuit import extract_circuit from validate_circuit import evaluate_consistency def main(): model load_or_train_model(epochs3) sparse_weights, meta decompose_model( model, energy_ratio0.9, sparsity_ratio0.7 ) for m in meta: print(flayer{m[name]}, shape{m[rank]}, keep_components{m[n_components]}) G extract_circuit(model, sparse_weights, top_edge_ratio0.2) print(fcircuit nodes: {G.number_of_nodes()}) print(fcircuit edges: {G.number_of_edges()}) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) test_dataset datasets.MNIST(./data, trainFalse, downloadTrue, transformtransform) loader DataLoader(test_dataset, batch_size256, shuffleFalse) cos evaluate_consistency(model, sparse_weights, loader, max_batches2) print(fcosine_similarity: {cos:.4f}) if __name__ __main__: main()预期的输出格式大致如下具体数值会因模型训练随机性而不同epoch1, avg_loss... epoch2, avg_loss... epoch3, avg_loss... layerfc1.weight, shapetorch.Size([256, 784]), keep_components... layerfc2.weight, shapetorch.Size([64, 256]), keep_components... layerfc3.weight, shapetorch.Size([10, 64]), keep_components... circuit nodes: ... circuit edges: ... cosine_similarity: 0.97xx如何判断实验结果是否符合预期circuit nodes应等于 784 256 64 10也就是所有输入神经元和隐藏层、输出层神经元总和circuit edges远小于全连接情况下的边数。全连接边数是784*256 256*64 64*10 217088条。稀疏化后如果仍是六位数说明sparsity_ratio或top_edge_ratio设置过宽如果不到五位数说明稀疏化有效cosine_similarity越高越好建议保持在 0.95 以上。如果你跑出来的cosine_similarity很低第一步不要调模型而是去看meta中每层保留了多少奇异值。如果某个隐藏层keep_components只剩 1 或 2基本可以断定这层被过压缩了需要提高energy_ratio。7. 常见问题与排查方法问题现象可能原因排查方式解决方案输出余弦相似度偏低energy_ratio太小低秩近似丢掉了关键信息打印每层保留的奇异值数量把energy_ratio从 0.9 提高到 0.95 或 0.99图边数仍然很大sparsity_ratio或top_edge_ratio设置偏保守统计每条边的权重绝对值分位数降低sparsity_ratio到 0.8 以上或降低top_edge_ratio图边数过少且不连通稀疏化做过头把中等强度边全部置零检查各层独立阈值不要用全局固定值按层统计权重分布用每层分位数作为阈值相同参数下每次结果不一致模型训练随机性导致权重分布不同固定随机种子记录实验参数在训练前设置torch.manual_seed并固定数据加载顺序torch.linalg.svd报维度错误传入的权重不是二维矩阵检查param.dim()只处理二维权重在代码里加param.dim() 2过滤条件输出层也被 ReLU 截断手动前向中激活函数应用位置错误检查网络结构输出层前通常不加激活参考forward_with_sparse_weights中的判断逻辑8. 最佳实践与工程建议8.1 参数不要一劳永逸先做能量分布分析不同层的最佳energy_ratio差异很大。输入层通常保留率高因为原始像素特征相对分散靠近输出层可以低一些因为类别信息集中。在日常工程里建议先写出每层奇异值平方的累积分布看一眼“前 80% 能量集中在前多少个奇异值”再决定每层的截断位置。8.2 稀疏化阈值按层计算不用全局百分比全局阈值的问题在于如果某一层权重整体数值偏小全局阈值会把这一层直接抹掉导致子网络断层。更稳妥的做法是逐层取分位数保证每层都有一定比例的边存活。8.3 把“稀疏权重分解”和“剪枝”分开看待剪枝的目的是减少计算量所以标准是“模型精度不掉”。电路提取的目的是可解释性,所以标准是“主要计算路径还在”。有时精度不掉但关键路径已经断了有时精度掉了几个点路径反而更清晰。明确你要解决什么问题再决定参数怎么调。8.4 生产环境要做行为回归验证如果你打算把提取出的电路用于规则审计或辅助决策建议在验证集上固定一组测试样本把“余弦相似度”和“标签一致率”同时纳入回归指标。每次调整稀疏参数后跑同一组样本方便横向对比。8.5 注意参数安全和权限边界如果你要分析的是生产模型尤其是涉及敏感业务或用户数据的模型务必遵守模型所在环境的访问控制规范只读取权重不修改线上配置。实验使用的数据也要满足最小必要原则先用脱敏数据验证流程再接触敏感样本。9. 总结与后续学习方向回到开头的问题电路提取难在权重预处理这是很多教程不会强调的细节。把稀疏权重分解作为前置步骤不仅让提取出的计算子图更清晰也让“为什么保留这条边”这个问题第一次有了可以解释的答案——因为这条边属于对整体能量贡献最大的主成分。如果你接下来想深入有三个方向值得关注把方法从 MLP 扩展到卷积网络和 Transformer重点研究“通道级分解”而不是单个权重分解把分解后的电路用于对抗样本分析观察不同类别样本是否走了不同子图把你的稀疏权重分解过程改成可微的放进训练过程中做端到端结构学习。建议你先跑通这篇的最小示例观察 SVD 奇异值分布和最终图结构的关系。理解了这个关系再切换到自己的模型上做参数调节会比盲目套模板扎实很多。