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

TensorFlow/Keras自定义训练循环:训练与验证阶段模式切换问询

解决自定义训练循环中训练/验证阶段的模式切换问题

嘿,这个问题确实是自定义训练循环里很容易踩的坑——毕竟Keras的fit()方法帮我们自动处理了训练/验证模式的切换,但自己写循环的时候就得手动管起来了。核心关键就是利用Keras模型调用时的training参数,来控制Dropout、BatchNormalization这类有模式差异的层的行为。

1. 训练阶段的正确处理

在训练循环里,你必须明确给模型调用加上training=True,这样:

  • Dropout层会随机失活神经元,实现正则化效果
  • BatchNormalization层会计算当前批次的均值和方差,同时更新全局的移动均值/方差

修改你训练阶段的代码如下:

for epoch in range(3):
    # 遍历训练数据集的批次。
    for step, (x_batch_train, y_batch_train) in enumerate(train_dataset):
        with tf.GradientTape() as tape:
            # 关键:加上training=True,强制模型进入训练模式
            logits1, logits2 = model(x_batch_train, training=True)
            loss_value1 = loss_fn1(y_batch_train[0], logits1)
            loss_value2 = loss_fn2(y_batch_train[1], logits2)
        grads1 = tape.gradient(loss_value1, model.trainable_weights[selection1])
        grads2 = tape.gradient(loss_value2, model.trainable_weights[selection2])
        optimizer1.apply_gradients(zip(grads1, model.trainable_weights[selection1]))
        optimizer2.apply_gradients(zip(grads2, model.trainable_weights[selection2]))

2. 验证阶段的正确处理

验证的时候,必须传递training=False,这样:

  • Dropout层会关闭,所有神经元都参与计算,避免验证结果被随机失活干扰
  • BatchNormalization层会使用训练阶段积累的全局移动均值/方差,而不是当前验证批次的统计量,保证验证结果的一致性

修改验证循环的代码如下:

# 在每个轮次结束后运行验证循环。
for x_batch_val, y_batch_val in val_dataset:
    # 关键:加上training=False,强制模型进入推理/验证模式
    val_logits1, val_logits2 = model(x_batch_val, training=False)
    # 接下来执行你的评估逻辑,比如计算准确率、验证损失等
    val_loss1 = loss_fn1(y_batch_val[0], val_logits1)
    val_loss2 = loss_fn2(y_batch_val[1], val_logits2)
    # ... 其他评估步骤(比如计算准确率、收集指标等)

额外注意:子类化模型的特殊处理

如果你是用子类化API定义的模型(比如class MyModel(tf.keras.Model)),那你的call方法必须接收training参数,并且把它传递给需要区分模式的层。比如:

class MyModel(tf.keras.Model):
    def __init__(self):
        super().__init__()
        self.dense1 = tf.keras.layers.Dense(64)
        self.dropout = tf.keras.layers.Dropout(0.5)
        self.bn = tf.keras.layers.BatchNormalization()
        self.dense2 = tf.keras.layers.Dense(10)
        self.dense3 = tf.keras.layers.Dense(20)

    def call(self, inputs, training=False):
        x = self.dense1(inputs)
        # 必须把training参数传递给Dropout和BatchNormalization
        x = self.dropout(x, training=training)
        x = self.bn(x, training=training)
        logits1 = self.dense2(x)
        logits2 = self.dense3(x)
        return logits1, logits2

如果你的子类化模型里没这么做,即使你在调用模型时传了training=True/False,那些层也不会响应模式切换,这也是很多人踩坑的点。

为什么自定义循环要手动处理?

因为Keras的model.fit()方法内部会自动在训练时设置training=True,在验证(validation_data)时设置training=False。但自定义训练循环完全由你掌控,所以必须手动指定这个参数,才能让Dropout、BN这些层表现出正确的行为。

内容的提问来源于stack exchange,提问作者Daniel Müller

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 07:41:22