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

如何在model.fit中使用不同的数据生成器训练CNN?

同时用两个ImageDataGenerator生成器训练CNN的两种方法

方法一:用tf.data.Dataset合并生成器

ImageDataGenerator的flow_from_directory返回的迭代器可以转成tf.data.Dataset,再把两个数据集合并,具体操作如下:

import tensorflow as tf

# 将train_gen1转换为tf.data.Dataset格式
ds1 = tf.data.Dataset.from_generator(
    lambda: train_gen1,
    output_types=(tf.float32, tf.float32),
    output_shapes=([None, 96, 96, 3], [None, 你的类别数])  # 替换为实际的类别数量
)

# 同理处理train_gen2
ds2 = tf.data.Dataset.from_generator(
    lambda: train_gen2,
    output_types=(tf.float32, tf.float32),
    output_shapes=([None, 96, 96, 3], [None, 你的类别数])
)

# 合并两个数据集,可选择concatenate直接合并或interleave交替取数据
combined_ds = ds1.concatenate(ds2).shuffle(1000).repeat()

# 训练时必须指定steps_per_epoch,计算方式为两个生成器总样本数除以批次大小
history = model.fit(
    x=combined_ds,
    epochs=epochs,
    validation_data=valid_gen,
    callbacks=noaug_callbacks,
    class_weight=class_weight,
    steps_per_epoch=(train_gen1.samples + train_gen2.samples) // train_gen1.batch_size
).history

方法二:自定义组合生成器

如果不想使用tf.data,可以直接写一个继承keras.utils.Sequence的类,手动合并两个生成器的输出:

from tensorflow.keras.utils import Sequence
import numpy as np

class CombinedGenerator(Sequence):
    def __init__(self, gen1, gen2):
        self.gen1 = gen1
        self.gen2 = gen2
        self.total_samples = gen1.samples + gen2.samples
        self.batch_size = gen1.batch_size  # 假设两个生成器批次大小一致,不一致需自行调整
    
    def __len__(self):
        return self.total_samples // self.batch_size
    
    def __getitem__(self, idx):
        # 从两个生成器各取一批数据
        x1, y1 = self.gen1[idx % len(self.gen1)]
        x2, y2 = self.gen2[idx % len(self.gen2)]
        # 合并当前批次的特征和标签
        x = np.concatenate([x1, x2], axis=0)
        y = np.concatenate([y1, y2], axis=0)
        # 可选:打乱当前批次的数据
        shuffle_idx = np.random.permutation(len(x))
        return x[shuffle_idx], y[shuffle_idx]
    
    def on_epoch_end(self):
        # 每个epoch结束后,让两个生成器各自刷新数据
        self.gen1.on_epoch_end()
        self.gen2.on_epoch_end()

# 创建组合生成器实例
combined_gen = CombinedGenerator(train_gen1, train_gen2)

# 直接使用组合生成器训练模型
history = model.fit(
    x=combined_gen,
    epochs=epochs,
    validation_data=valid_gen,
    callbacks=noaug_callbacks,
    class_weight=class_weight
).history

注意事项

  • 如果两个生成器的批次大小不一致,自定义生成器中需要做适配调整,或者在tf.data流程中用batch()重新设置统一批次。
  • 使用tf.data方式时必须设置steps_per_epoch,否则Keras无法确定每个epoch需要执行多少步。
  • 若两个生成器是对同一数据集做不同数据增强,合并后相当于扩充训练数据;若对应不同数据集,则是混合训练两类数据。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 05:35:50