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

如何获取PyTorch神经网络中各神经元的输入与输出边权重?

实现方案

核心逻辑说明

  • 神经元用(层名, 神经元下标)作为唯一标识,比如('fc2', 1)代表第2个全连接层的第2个神经元(下标从0开始计数)
  • PyTorch全连接层nn.Linear的权重weight形状固定为(输出神经元数, 输入神经元数),因此:
    • 单个神经元的入连接:对应所在层weight矩阵当前神经元下标对应的行,每行的每个元素就是上一层对应神经元连到当前神经元的权重
    • 单个神经元的出连接:对应下一层weight矩阵当前神经元下标对应的列,每列的每个元素就是当前神经元连到下一层对应神经元的权重
    • 输入层神经元仅存在出连接,最后一层神经元仅存在入连接
      注意:如果你的网络包含非全连接层,只需调整fc_layers的过滤逻辑即可,本实现仅适配你给出的全连接串行结构

完整实现代码

import torch
import random
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):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        x = self.fc3(x)
        return x

net = Model()

# 按顺序存储所有全连接层的基础信息
fc_layers = []
for name, module in net.named_modules():
    if isinstance(module, nn.Linear):
        fc_layers.append({
            "name": name,
            "module": module,
            "in_dim": module.in_features,
            "out_dim": module.out_features
        })

# 生成全网络所有神经元的列表,标识格式为(层名, 神经元下标)
all_neurons = []
# 先添加输入层神经元
for i in range(fc_layers[0]['in_dim']):
    all_neurons.append(('input', i))
# 添加各全连接层的输出神经元
for layer in fc_layers:
    for i in range(layer['out_dim']):
        all_neurons.append((layer['name'], i))

def get_in_connections(neuron):
    layer_name, neuron_idx = neuron
    # 输入层没有入连接
    if layer_name == 'input':
        return []
    # 匹配当前层信息
    layer_info = next(l for l in fc_layers if l['name'] == layer_name)
    in_edges = []
    prev_layer_name = 'input' if fc_layers.index(layer_info) == 0 else fc_layers[fc_layers.index(layer_info)-1]['name']
    for prev_neuron_idx in range(layer_info['in_dim']):
        in_edges.append({
            "from_neuron": (prev_layer_name, prev_neuron_idx),
            "to_neuron": neuron,
            "weight": layer_info['module'].weight[neuron_idx, prev_neuron_idx],
            # 存储权重在原矩阵的位置引用,方便后续直接修改原模型权重
            "weight_ref": (layer_info['module'].weight, neuron_idx, prev_neuron_idx)
        })
    return in_edges

def get_out_connections(neuron):
    layer_name, neuron_idx = neuron
    # 最后一层没有出连接
    last_layer_name = fc_layers[-1]['name']
    if layer_name == last_layer_name:
        return []
    # 匹配下一层信息
    if layer_name == 'input':
        next_layer_idx = 0
    else:
        current_layer_idx = next(i for i, l in enumerate(fc_layers) if l['name'] == layer_name)
        next_layer_idx = current_layer_idx + 1
    next_layer_info = fc_layers[next_layer_idx]
    out_edges = []
    for next_neuron_idx in range(next_layer_info['out_dim']):
        out_edges.append({
            "from_neuron": neuron,
            "to_neuron": (next_layer_info['name'], next_neuron_idx),
            "weight": next_layer_info['module'].weight[next_neuron_idx, neuron_idx],
            "weight_ref": (next_layer_info['module'].weight, next_neuron_idx, neuron_idx)
        })
    return out_edges

def get_neuron_in_out_edges(list_of_neurons):
    # 打乱输入的神经元列表
    shuffled_list_of_neurons = list_of_neurons.copy()
    random.shuffle(shuffled_list_of_neurons)

    in_connections_list = []
    out_connections_list = []

    for neuron in shuffled_list_of_neurons:
        in_connections = get_in_connections(neuron)
        out_connections = get_out_connections(neuron)
        in_connections_list.append([neuron, in_connections])
        out_connections_list.append([neuron, out_connections])
    
    return in_connections_list, out_connections_list

使用示例

1. 随机选择3个神经元获取出入连接

selected_neurons = random.sample(all_neurons, 3)
in_edges_list, out_edges_list = get_neuron_in_out_edges(selected_neurons)

2. 把获取到的权重中小于10的数值修改为10

# 封装修改逻辑避免重复代码
def modify_weights(edges_list):
    for neuron_item in edges_list:
        edges = neuron_item[1]
        for edge in edges:
            weight_mat, row, col = edge['weight_ref']
            if weight_mat[row, col] < 10:
                # 加torch.no_grad避免修改被计入计算图
                with torch.no_grad():
                    weight_mat[row, col] = 10

modify_weights(in_edges_list)
modify_weights(out_edges_list)

内容的提问来源于stack exchange,提问作者Penguin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 20:06:04