Keras中自定义初始值的可学习权重矩阵乘法问题排查
哦,我明白你的问题了——你创建的tf.Variable确实存在,但Keras根本没把它当作模型的可训练权重来追踪,所以训练时完全不会更新它。问题出在你直接用K.dot操作这个变量,没有把它封装到Keras的层结构里,Keras的模型只会管理属于层的参数。
给你两个实用的解决办法,先讲最简单的那种:
方法1:用Dense层(最省事的方案)
Dense层本质就是做矩阵乘法(默认加偏置,我们可以把偏置关掉),刚好匹配你的需求。只需要创建一个不带偏置的Dense层,再手动把传入的matrix设为它的初始权重就行:
def create_model(num_columns, matrix): inp_layer = tfl.Input((num_columns,)) dense = tfl.Dense(512, activation='relu')(inp_layer) dense = tfl.Dense(256, activation='relu')(dense) dense = tfl.Dense(128, activation='relu')(dense) # 创建不带偏置的Dense层,输出维度对应matrix的列数(206) dense = tfl.Dense(matrix.shape[1], use_bias=False)(dense) # 将传入的matrix转为float32后设为该层的初始权重 dense.set_weights([matrix.astype(np.float32)]) model = tf.keras.Model(inputs=inp_layer, outputs=dense) model.compile(optimizer='adam', loss=['binary_crossentropy']) model.summary() return model matrix = np.random.randint(0,2,(128, 206)) # 实际为有意义的数值,非随机 num_columns = 750 model = create_model(num_columns,matrix)
运行这段代码后,你会在model.summary()里看到这个Dense层有128*206=26368个可训练参数,就是你传入的矩阵,训练时会正常被微调。
方法2:自定义可训练层(适合复杂场景)
如果之后你需要对这个矩阵做更复杂的操作(比如加正则、自定义更新逻辑),可以自定义一个Keras层来明确管理这个可训练参数:
class TrainableMatrixMultiply(tfl.Layer): def __init__(self, initial_matrix, **kwargs): self.initial_matrix = initial_matrix super().__init__(**kwargs) def build(self, input_shape): # 在build方法里注册可训练变量,Keras会自动追踪它 self.matrix = self.add_weight( name='trainable_matrix', shape=self.initial_matrix.shape, initializer=tfl.initializers.Constant(self.initial_matrix), trainable=True ) super().build(input_shape) def call(self, inputs): # 执行矩阵乘法逻辑 return K.dot(inputs, self.matrix) def create_model(num_columns, matrix): inp_layer = tfl.Input((num_columns,)) dense = tfl.Dense(512, activation='relu')(inp_layer) dense = tfl.Dense(256, activation='relu')(dense) dense = tfl.Dense(128, activation='relu')(dense) # 使用自定义层传入初始矩阵 dense = TrainableMatrixMultiply(matrix)(dense) model = tf.keras.Model(inputs=inp_layer, outputs=dense) model.compile(optimizer='adam', loss=['binary_crossentropy']) model.summary() return model
这个自定义层会在build阶段把矩阵注册为可训练权重,Keras会自动把它加入模型的参数列表,summary里也能看到对应的可训练参数项。
为什么你的原代码不行?
你直接在函数里创建tf.Variable然后用K.dot,这个变量不属于任何Keras层——Keras的模型只会追踪那些通过层的add_weight方法创建的变量(或是内置层自带的参数)。所以这个变量虽然存在,但模型根本不把它当作自己的可训练参数,训练时自然不会更新它。
内容的提问来源于stack exchange,提问作者Antonio Paladini
相关产品推荐
相关产品推荐

