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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 22:29:54