You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效获取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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 00:45:04