如何将Keras层中特定的权重分量设置为不可训练?
Keras 权重矩阵部分分量设置为不可训练的实现方案
原生Keras的全连接层默认将整个权重矩阵作为单个可训练变量管理,仅支持对完整变量设置可训练状态,若要实现部分分量固定,可采用以下两种成熟方案:
方案1:权重拆分拼接法(适合整行/整列固定场景)
直接将需要固定的部分和可训练的部分拆分为两个独立变量,其中固定部分设置为不可训练,使用时拼接为完整权重矩阵即可,刚好适配你提出的W1、W2拆分需求。
代码实现
import tensorflow as tf from tensorflow import keras class PartialFixedDense(keras.layers.Layer): def __init__(self, units, fixed_row_count=1, **kwargs): super().__init__(**kwargs) self.units = units # 前fixed_row_count行属于不可训练的W1 self.fixed_row_count = fixed_row_count def build(self, input_shape): input_dim = input_shape[-1] # 定义不可训练的W1 self.W1 = self.add_weight( shape=(self.fixed_row_count, self.units), initializer="glorot_uniform", trainable=False, name="fixed_weight" ) # 定义可训练的W2 self.W2 = self.add_weight( shape=(input_dim - self.fixed_row_count, self.units), initializer="glorot_uniform", trainable=True, name="trainable_weight" ) # 拼接为完整权重矩阵 self.full_W = tf.concat([self.W1, self.W2], axis=0) # 偏置根据需求设置可训练状态 self.b = self.add_weight( shape=(self.units,), initializer="zeros", trainable=True, name="bias" ) def call(self, inputs): return tf.matmul(inputs, self.full_W) + self.b
使用示例
你举的2维输入、2维输出、第一行W1不可训练的场景,直接替换原有Dense层即可:
# 替换原有keras.layers.Dense(2) layer = PartialFixedDense(units=2, fixed_row_count=1)
方案2:梯度掩码法(适合任意位置固定场景)
如果需要固定的是零散的任意位置,不是连续的整行整列,可以通过梯度掩码实现:定义和权重形状相同的掩码矩阵,需要固定的位置标记为0,可训练的位置标记为1,反向传播时截断固定位置的梯度即可。
代码实现
class MaskedDense(keras.layers.Layer): def __init__(self, units, weight_mask, **kwargs): super().__init__(**kwargs) self.units = units # 掩码矩阵:固定位置为0,可训练位置为1 self.weight_mask = tf.constant(weight_mask, dtype=tf.float32) def build(self, input_shape): self.W = self.add_weight( shape=(input_shape[-1], self.units), initializer="glorot_uniform", trainable=True, name="weight" ) self.b = self.add_weight( shape=(self.units,), initializer="zeros", trainable=True, name="bias" ) def call(self, inputs): # 正向传播使用完整权重,反向传播仅更新掩码为1的位置 masked_W = self.W * self.weight_mask + tf.stop_gradient(self.W * (1 - self.weight_mask)) return tf.matmul(inputs, masked_W) + self.b
使用示例
比如要固定2×2权重矩阵的第一行,只需传入对应掩码即可:
# 第一行固定(标记为0),第二行可训练(标记为1) mask = [[0, 0], [1, 1]] layer = MaskedDense(units=2, weight_mask=mask)
内容的提问来源于stack exchange,提问作者narutoArea51
相关产品推荐
相关产品推荐

