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

PyTorch中如何固定神经网络某层部分神经元的权重与偏置?

实现思路与代码示例

不需要使用LocallyConnected1D,直接基于PyTorch的普通全连接层(nn.Linear)就能实现需求,以下是两种实用方案:

方案一:单全连接层拆分参数控制梯度

直接用一个总神经元数为X=X1+X2的nn.Linear层,手动设置前X1个神经元的权重、偏置不参与梯度更新,剩下X2个保持可训练:

代码示例

import torch
import torch.nn as nn

class CustomNet(nn.Module):
    def __init__(self, in_features, total_neurons, fixed_neurons_num):
        super().__init__()
        # 定义包含所有神经元的全连接层
        self.target_layer = nn.Linear(in_features, total_neurons)
        
        # 固定前fixed_neurons_num个神经元的权重与偏置
        # 权重形状:(total_neurons, in_features),前fixed_neurons_num行对应固定神经元
        self.target_layer.weight[:fixed_neurons_num].requires_grad = False
        # 偏置形状:(total_neurons,),前fixed_neurons_num个元素对应固定神经元
        self.target_layer.bias[:fixed_neurons_num].requires_grad = False
        
        # 后续网络层根据需求定义
        self.relu = nn.ReLU()
        self.output = nn.Linear(total_neurons, 10)

    def forward(self, x):
        x = self.target_layer(x)
        x = self.relu(x)
        x = self.output(x)
        return x

方案二:拆分为固定层+可训练层

把目标层拆成两个独立的nn.Linear层:一个包含X1个固定神经元(禁用梯度),另一个包含X2个可训练神经元,前向传播时拼接两个层的输出:

代码示例

import torch
import torch.nn as nn

class CustomNet(nn.Module):
    def __init__(self, in_features, fixed_neurons_num, trainable_neurons_num):
        super().__init__()
        # 固定神经元层:禁用梯度更新
        self.fixed_layer = nn.Linear(in_features, fixed_neurons_num)
        for param in self.fixed_layer.parameters():
            param.requires_grad = False
        
        # 可训练神经元层:正常参与训练
        self.trainable_layer = nn.Linear(in_features, trainable_neurons_num)
        
        # 后续网络层
        self.relu = nn.ReLU()
        self.output = nn.Linear(fixed_neurons_num + trainable_neurons_num, 10)

    def forward(self, x):
        fixed_out = self.fixed_layer(x)
        trainable_out = self.trainable_layer(x)
        # 拼接两个层的输出,得到总神经元的输出结果
        x = torch.cat([fixed_out, trainable_out], dim=1)
        x = self.relu(x)
        x = self.output(x)
        return x

关键注意事项

  1. 参数初始化:如果需要给固定神经元设置特定的初始权重/偏置,要在设置requires_grad=False之前完成初始化操作。
  2. 优化器适配:PyTorch的优化器默认只会更新requires_grad=True的参数,无需额外过滤,训练时直接传入模型参数即可。
  3. 可视化适配:方案一的网络结构在可视化时会显示为单个全连接层,更符合你的需求;方案二会显示两个独立层,结构更直观但可视化结果是拆分状态。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 18:34:52