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

使用Keras train_on_batch训练含Dropout层模型时,如何设置学习阶段?

解决Keras中Dropout训练/测试阶段切换的问题

首先,我注意到你的代码里有一个关键问题:discriminator是一个方法,每次调用它都会重新构建一个新的模型实例。如果你直接用self.discriminator.train_on_batch,相当于每次训练都在从头训练一个新模型,之前的权重不会被保留,这肯定不是你想要的。所以第一步要先修正这个问题:在类的初始化方法中创建一次判别器模型并保存为实例变量,后续训练和测试都复用这个模型。

接下来针对Dropout的阶段切换需求,分两种情况给出解决方案:

方案一:使用K.set_learning_phase(适用于旧版Keras/TensorFlow 1.x)

K.set_learning_phase可以全局设置当前的学习阶段:1代表训练模式(启用Dropout、BatchNorm的训练行为),0代表测试模式(禁用Dropout、使用BatchNorm的移动均值/方差)。

修改后的代码示例:

首先在类的__init__中初始化模型:

from keras.models import Model, Sequential
from keras.layers import Input, concatenate, Dropout, LSTM, Flatten, Dense
from keras.constraints import unit_norm
import keras.backend as K

class YourModelClass:
    def __init__(self, input_shape):
        self.shape = input_shape
        # 确保这个形状和concatenate后的输入匹配,比如(input_shape[0], input_shape[1]*2)
        self.shape_double = (input_shape[0], input_shape[1]*2)
        # 初始化判别器模型并保存为实例变量
        self.d_model = self.discriminator()

    def discriminator(self):
        x_A = Input(shape=self.shape)
        x_B = Input(shape=self.shape)
        x = concatenate([x_A, x_B], axis=-1)
        model = Sequential()
        model.add(Dropout(0.5, input_shape=self.shape_double))
        model.add(LSTM(200, return_sequences=True, kernel_constraint=unit_norm()))
        model.add(Dropout(0.5))
        model.add(LSTM(200, return_sequences=True, kernel_constraint=unit_norm()))
        model.add(Dropout(0.5))
        model.add(Flatten())
        model.add(Dense(8, activation="softmax", kernel_constraint=unit_norm()))
        label = model(x)
        return Model([x_A, x_B], label)

然后修改训练方法,设置训练阶段:

def train(self, epochs, batch_size):
    # 切换到训练模式,启用Dropout
    K.set_learning_phase(1)
    for epoch in range(epochs):
        total_loss = 0.0
        for batch, train_A, train_B, train_label in enumerate(Load_train(batch_size)):
            d_loss = self.d_model.train_on_batch([train_A, train_B], train_label)
            total_loss += d_loss
            # 可选:打印批次训练信息
            if batch % 10 == 0:
                print(f"Epoch {epoch+1}, Batch {batch}, D Loss: {d_loss:.4f}")
        print(f"Epoch {epoch+1} completed, Avg D Loss: {total_loss/(batch+1):.4f}")
    # 训练结束后切换回测试模式,避免影响后续测试
    K.set_learning_phase(0)

测试方法中确保处于测试模式:

def test(self, test_A, test_B, test_label):
    # 明确切换到测试模式,禁用Dropout
    K.set_learning_phase(0)
    predicted_label_dist = self.d_model.predict([test_A, test_B])
    # 这里可以添加准确率计算等逻辑
    predicted_labels = predicted_label_dist.argmax(axis=1)
    accuracy = (predicted_labels == test_label.argmax(axis=1)).mean()
    print(f"Test Accuracy: {accuracy:.4f}")
    return predicted_label_dist

方案二:使用training参数(推荐,适用于TensorFlow 2.x/Keras 2.3+)

在TensorFlow 2.x中,K.set_learning_phase已经被弃用,更推荐的方式是在调用模型时显式指定training参数,这样不会影响全局状态,更灵活。

修改后的训练和测试方法:

训练时,调用train_on_batch时传入training=True:

def train(self, epochs, batch_size):
    for epoch in range(epochs):
        total_loss = 0.0
        for batch, train_A, train_B, train_label in enumerate(Load_train(batch_size)):
            # 显式指定training=True,启用Dropout
            d_loss = self.d_model.train_on_batch([train_A, train_B], train_label, training=True)
            total_loss += d_loss
            if batch % 10 == 0:
                print(f"Epoch {epoch+1}, Batch {batch}, D Loss: {d_loss:.4f}")
        print(f"Epoch {epoch+1} completed, Avg D Loss: {total_loss/(batch+1):.4f}")

测试时,predict方法默认training=False,所以直接调用即可,也可以显式指定:

def test(self, test_A, test_B, test_label):
    # predict默认禁用Dropout,也可以显式指定training=False
    predicted_label_dist = self.d_model.predict([test_A, test_B])
    # 或者用模型直接调用的方式:
    # predicted_label_dist = self.d_model([test_A, test_B], training=False).numpy()
    predicted_labels = predicted_label_dist.argmax(axis=1)
    accuracy = (predicted_labels == test_label.argmax(axis=1)).mean()
    print(f"Test Accuracy: {accuracy:.4f}")
    return predicted_label_dist

额外注意事项

  • 确保self.shape_double和concatenate([x_A, x_B], axis=-1)后的张量形状完全匹配,否则会出现输入形状不兼容的错误。比如如果输入是(time_steps, features),拼接后是(time_steps, features*2),那么self.shape_double应该设为这个值。
  • 如果你使用的是TensorFlow 2.x,建议优先使用方案二,因为全局的学习阶段设置可能会和其他组件(比如混合精度训练、自定义层)产生冲突。

内容的提问来源于stack exchange,提问作者Yutas

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:55:30