如何获取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
相关产品推荐
相关产品推荐

