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

如何在Keras中同时训练两类回归数据?实现交错迭代训练

实现交错批次的混合训练方案(Keras 3.2.1 + Jax后端)

针对你的需求,直接用fit方法无法严格实现监督/自监督批次交替的逻辑,最灵活的方式是通过自定义训练循环结合tf.data.Dataset的交错处理来实现。以下是具体步骤和代码示例:

核心思路

  1. 分别构建监督数据和自监督数据的数据集
  2. 通过数据集操作实现两类批次的交错输出
  3. 自定义训练步骤,根据批次类型计算对应损失并更新模型参数

具体实现

1. 导入依赖并构建回归模型

import keras
from keras import layers
import jax.numpy as jnp
import tensorflow as tf  # Keras 3跨后端兼容,用于构建数据集
def build_regression_model(input_dim):
    inputs = layers.Input(shape=(input_dim,))
    x = layers.Dense(64, activation='relu')(inputs)
    x = layers.Dense(32, activation='relu')(x)
    outputs = layers.Dense(1)(x)
    return keras.Model(inputs=inputs, outputs=outputs)

# 假设输入向量维度为10,根据实际情况调整
model = build_regression_model(input_dim=10)

2. 准备示例数据(替换为你的真实数据)

# 监督数据:输入向量 + 对应标签
x_supervised = jnp.random.rand(1000, 10)
y_supervised = jnp.random.rand(1000, 1)

# 自监督数据:需要预测结果相似的输入向量对
x1_self = jnp.random.rand(1000, 10)
x2_self = jnp.random.rand(1000, 10)

batch_size = 32

3. 构建并交错数据集

通过zip和flat_map实现监督批次与自监督批次的交替输出:

# 构建监督数据集:打乱、分批、预取
supervised_ds = tf.data.Dataset.from_tensor_slices((x_supervised, y_supervised))
supervised_ds = supervised_ds.shuffle(1000).batch(batch_size).prefetch(tf.data.AUTOTUNE)

# 构建自监督数据集:打乱、分批、预取
self_supervised_ds = tf.data.Dataset.from_tensor_slices((x1_self, x2_self))
self_supervised_ds = self_supervised_ds.shuffle(1000).batch(batch_size).prefetch(tf.data.AUTOTUNE)

# 交错两个数据集:监督批次 → 自监督批次 → 监督批次...
interleaved_ds = tf.data.Dataset.zip((supervised_ds, self_supervised_ds))
interleaved_ds = interleaved_ds.flat_map(
    lambda s_batch, ss_batch: tf.data.Dataset.from_tensors(s_batch).concatenate(
        tf.data.Dataset.from_tensors(ss_batch)
    )
)

# 若两类数据样本量不同,给样本少的数据集加repeat()避免提前耗尽
# self_supervised_ds = self_supervised_ds.repeat()

4. 定义优化器与训练步骤

自定义训练逻辑,根据批次结构区分两类数据并计算对应损失:

optimizer = keras.optimizers.Adam(learning_rate=1e-3)
mse_loss = keras.losses.MeanSquaredError()

@tf.function  # 编译为计算图,提升Jax后端下的训练效率
def train_step(batch):
    # 判断批次类型:监督批次是(input, label)二元组,自监督批次是(input1, input2)二元组
    x, y_or_x2 = batch
    if y_or_x2.shape[-1] == 1:
        # 处理监督数据:计算预测值与标签的MSE
        with tf.GradientTape() as tape:
            y_pred = model(x, training=True)
            loss = mse_loss(y_or_x2, y_pred)
    else:
        # 处理自监督数据:计算两组预测值的MSE
        with tf.GradientTape() as tape:
            y1_pred = model(x, training=True)
            y2_pred = model(y_or_x2, training=True)
            loss = mse_loss(y1_pred, y2_pred)
    
    # 计算梯度并更新模型参数
    gradients = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(gradients, model.trainable_variables))
    return loss

5. 执行训练循环

epochs = 10
for epoch in range(epochs):
    print(f"Epoch {epoch+1}/{epochs}")
    total_loss = 0.0
    batch_count = 0
    
    for batch in interleaved_ds:
        loss = train_step(batch)
        total_loss += loss.numpy()
        batch_count += 1
        
        # 每10个批次打印一次实时损失
        if batch_count % 10 == 0:
            print(f"Batch {batch_count}, Current Loss: {loss.numpy():.4f}")
    
    avg_epoch_loss = total_loss / batch_count
    print(f"Epoch Average Loss: {avg_epoch_loss:.4f}\n")

关键说明

  • 数据集交错逻辑:通过zip将两类批次打包,再用flat_map拆分为交替的单个批次,严格保证训练时的批次顺序。
  • 训练步骤兼容性:@tf.function装饰器在Jax后端下会被Keras自动转换为Jax计算图,不影响性能。
  • 灵活调整:如果需要调整两类数据的训练占比(比如每2个监督批次对应1个自监督批次),可以修改flat_map中的数据集拼接逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 20:20:57