如何高效获取PyTorch神经网络中的neuron-edge-neuron数值
实现思路
- 用前向钩子自动捕获所有全连接层的输入、输出,无需手动修改网络
forward函数返回值,可适配任意包含nn.Linear层的网络 - 全程基于PyTorch原生张量广播机制做节点、权重匹配,无Python层循环,性能足以支撑大参数量网络,支持GPU加速
- 自动兼容第一层输入的特殊逻辑,直接将网络输入作为第一层源节点值,和后续层处理逻辑统一
代码实现
import torch import torch.nn as nn class Model(nn.Module): def __init__(self): super(Model, self).__init__() self.fc1 = nn.Linear(1, 2) self.fc2 = nn.Linear(2, 3) self.fc3 = nn.Linear(3, 1) def forward(self, x): x1 = self.fc1(x) x = torch.relu(x1) x2 = self.fc2(x) x = torch.relu(x2) x3 = self.fc3(x) return x3, x2, x1 net = Model() # 存储全连接层的输入值、输出值、权重参数 layer_io = [] def hook_fn(module, input, output): if isinstance(module, nn.Linear): # 单样本场景去掉batch维度,多样本场景可按需调整维度处理逻辑 layer_io.append((input[0].squeeze(0), output.squeeze(0), module.weight)) # 给所有全连接层注册前向钩子 for module in net.modules(): if isinstance(module, nn.Linear): module.register_forward_hook(hook_fn) # 传入测试输入,前向传播自动捕获各层数据 x = torch.randn(1, 1) net(x) # 生成 [源节点值, 边权重, 目标节点值] 三元组 all_triplets = [] for src_nodes, dst_nodes, weight in layer_io: # 张量广播匹配所有节点组合,无需循环 src_expand = src_nodes.repeat_interleave(len(dst_nodes)) weight_flat = weight.flatten() dst_expand = dst_nodes.repeat(len(src_nodes)) # 堆叠为该层所有三元组 layer_triplets = torch.stack([src_expand, weight_flat, dst_expand], dim=1) all_triplets.append(layer_triplets) # 验证输出 for idx, triplet in enumerate(all_triplets): print(f"第{idx+1}层三元组形状:{triplet.shape}")
输出说明
你提供的示例网络输出完全符合需求:
- 第1层(输入到fc1):形状为
[2, 3],共2组三元组,对应1个输入节点 * 2个输出神经元 - 第2层(fc1到fc2):形状为
[6, 3],共6组三元组,对应2个输入神经元 * 3个输出神经元 - 第3层(fc2到fc3):形状为
[3, 3],共3组三元组,对应3个输入神经元 * 1个输出神经元
如果需要转成嵌套列表格式,直接对张量调用.tolist()方法即可。
性能优化提示
- 大参数量网络下可以直接基于张量做可视化预处理,无需转Python列表,可大幅降低内存占用
- 多batch场景下可以在钩子中对batch维度做平均,或者保留每个样本的独立三元组,按需调整维度处理逻辑即可
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

