如何在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
相关产品推荐
相关产品推荐

