如何在TensorFlow 2.x每个Epoch开始时更新不可训练变量?
TensorFlow 2.x中回调更新不可训练层变量不生效的解决方法
问题描述
在TensorFlow 2.x中尝试通过回调更新自定义层的不可训练变量weighted_add_layer.weight,但模型编译后,fit始终使用编译时的初始值,无法后续更新。试过tf.keras.backend.set_value()等方法都无效,演示代码如下:
class WeightedAddLayer(tf.keras.layers.Layer): def __init__(self, weight=0.00, *args, **kwargs): super(WeightedAddLayer, self).__init__(*args, **kwargs) self.weight = tf.Variable(0., trainable=False) def add(self, inputA, inputB): return (self.weight * inputA + self.weight * inputB) def update(self, weight): tf.keras.backend.set_value(self.weight, weight) input_A = tfkl.Input( shape=(32), batch_size=32, ) input_B = tfkl.Input( shape=(32), batch_size=32, ) weighted_add_layer = WeightedAddLayer() output = weighted_add_layer.add(input_A, input_B) model = tfk.Model( inputs=[input_A, input_B], outputs=[output], ) model.compile( optimizer='adam', loss=losses.MeanSquaredError() ) # Custom callback function def update_fun(epoch, steps=50): weighted_add_layer.update( tf.clip_by_value( epoch / steps, clip_value_min=tf.constant(0.0), clip_value_max=tf.constant(1.0),) ) # Custom callback update_callback = tfk.callbacks.LambdaCallback( on_epoch_begin=lambda epoch, logs: update_fun(epoch) ) # train model history = model.fit( x=train_data, epochs=EPOCHS, validation_data=valid_data, callbacks=[update_callback], )
解决建议
问题根源:TensorFlow 2.x在模型编译后会固化计算图,
tf.keras.backend.set_value()属于Python端的赋值操作,不会触发计算图的重新追踪,导致训练时仍然使用初始值。方案1:改用变量的
.assign()方法
修改自定义层的update方法,直接使用TensorFlow变量的.assign()方法,该操作会被纳入计算图追踪,确保训练时使用更新后的值:
def update(self, weight): self.weight.assign(weight)
- 方案2:先设为可训练再冻结(可选)
如果需要变量被计算图追踪但不参与训练,可以先将变量设为trainable=True,初始化层后再冻结它:
class WeightedAddLayer(tf.keras.layers.Layer): def __init__(self, weight=0.00, *args, **kwargs): super(WeightedAddLayer, self).__init__(*args, **kwargs) self.weight = tf.Variable(0., trainable=True) # 初始设为可训练 # 初始化层后冻结变量 weighted_add_layer.weight.trainable = False
之后同样用.assign()方法更新变量即可。
- 方案3:在回调中通过模型实例获取层
如果担心直接引用层实例存在作用域问题,可以在回调中通过model.layers遍历找到目标层再更新:
def update_fun(epoch, model, steps=50): # 遍历模型层找到自定义层 weighted_add_layer = next(layer for layer in model.layers if isinstance(layer, WeightedAddLayer)) new_weight = tf.clip_by_value(epoch / steps, 0.0, 1.0) weighted_add_layer.weight.assign(new_weight) # 回调中传入model实例 update_callback = tfk.callbacks.LambdaCallback( on_epoch_begin=lambda epoch, logs: update_fun(epoch, model) )
内容的提问来源于stack exchange,提问作者Anirban Mukherjee
相关产品推荐
相关产品推荐

