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

如何在TensorFlow/PyTorch中部分冻结神经网络层参数?

全连接层部分参数固定为0并仅训练剩余参数的实现方案

针对你需要固定全连接层中约800个参数为0、仅训练剩余130个参数的需求,主流深度学习框架都有成熟的实现方式,以下是PyTorch和TensorFlow的具体操作步骤:

PyTorch 实现方式

步骤1:初始化层并生成掩码

首先定义全连接层,然后创建一个与权重同形状的掩码——需要训练的参数位置设为1,固定为0的位置设为0,同时初始将权重按掩码置0:

import torch
import torch.nn as nn

# 定义30->30的全连接层
fc_layer = nn.Linear(30, 30)

# 生成掩码:随机选择130个参数位置设为1(若有固定位置需求,可手动指定索引)
mask = torch.zeros_like(fc_layer.weight)
# 展平权重后随机选取130个索引
indices = torch.randperm(mask.numel())[:130]
mask.view(-1)[indices] = 1.0

# 初始将不需要训练的参数置0
fc_layer.weight.data *= mask

步骤2:约束参数更新

有两种高效方式确保固定参数始终为0:

方式一:注册梯度钩子

通过钩子过滤梯度,让固定参数的梯度为0,优化器不会更新这些参数:

def mask_grad(grad):
    # 仅保留掩码为1位置的梯度
    return grad * mask

# 给权重注册梯度钩子
fc_layer.weight.register_hook(mask_grad)

方式二:训练循环后置0

每次优化器更新参数后,手动用掩码将固定参数重置为0:

# 假设已定义optimizer和dataloader
optimizer = torch.optim.Adam(fc_layer.parameters(), lr=1e-3)

for x, y in dataloader:
    optimizer.zero_grad()
    output = fc_layer(x)
    loss = ... # 定义损失函数
    loss.backward()
    optimizer.step()
    
    # 重置固定参数为0
    with torch.no_grad():
        fc_layer.weight.data *= mask

TensorFlow/Keras 实现方式

方式一:自定义掩码全连接层

自定义层在每次前向传播时约束参数,确保固定位置始终为0:

import tensorflow as tf
from tensorflow.keras.layers import Layer

class MaskedDense(Layer):
    def __init__(self, units, mask, use_bias=True, **kwargs):
        super().__init__(**kwargs)
        self.units = units
        self.mask = mask
        self.use_bias = use_bias

    def build(self, input_shape):
        # 初始化权重
        self.kernel = self.add_weight(
            shape=(input_shape[-1], self.units),
            initializer='glorot_uniform',
            trainable=True
        )
        # 初始按掩码置0
        self.kernel.assign(self.kernel * self.mask)
        
        if self.use_bias:
            self.bias = self.add_weight(
                shape=(self.units,),
                initializer='zeros',
                trainable=True
            )
        super().build(input_shape)

    def call(self, inputs):
        # 前向传播时确保参数被掩码约束
        self.kernel.assign(self.kernel * self.mask)
        output = tf.matmul(inputs, self.kernel)
        if self.use_bias:
            output += self.bias
        return output

# 生成掩码(若有固定位置需求,手动指定索引即可)
mask = tf.zeros((30, 30))
indices = tf.random.shuffle(tf.range(30*30))[:130]
mask = tf.tensor_scatter_nd_update(mask, tf.reshape(indices, (-1, 1)), tf.ones(130))
mask = tf.reshape(mask, (30, 30))

# 使用自定义层构建模型
model = tf.keras.Sequential([
    MaskedDense(30, mask=mask, input_shape=(30,))
])

方式二:训练回调约束参数

使用普通Dense层,通过回调函数在每次训练批次后重置固定参数:

# 定义回调类
class MaskCallback(tf.keras.callbacks.Callback):
    def __init__(self, mask, target_layer):
        super().__init__()
        self.mask = mask
        self.target_layer = target_layer

    def on_train_batch_end(self, batch, logs=None):
        # 重置固定参数为0
        self.target_layer.kernel.assign(self.target_layer.kernel * self.mask)

# 初始化普通全连接层
fc_layer = tf.keras.layers.Dense(30, input_shape=(30,))
model = tf.keras.Sequential([fc_layer])

# 生成并应用初始掩码
mask = tf.zeros((30, 30))
indices = tf.random.shuffle(tf.range(30*30))[:130]
mask = tf.tensor_scatter_nd_update(mask, tf.reshape(indices, (-1, 1)), tf.ones(130))
mask = tf.reshape(mask, (30, 30))
fc_layer.kernel.assign(fc_layer.kernel * mask)

# 编译并训练模型
model.compile(optimizer='adam', loss='mse')
model.fit(x_train, y_train, callbacks=[MaskCallback(mask, fc_layer)])

注意事项

  • 若需要固定部分偏置参数,可参照权重的掩码逻辑处理偏置项;
  • 掩码的索引可根据业务需求手动指定,无需随机选择;
  • PyTorch的梯度钩子方式比后置0更高效,因为直接过滤了梯度,避免无效的参数更新计算。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 10:05:22