如何在MirroredStrategy下结合ImageDataGenerator.flow_from_directory使用model.fit()?
解决MirroredStrategy配合ImageDataGenerator转tf.data.Dataset的报错问题
我帮你分析下报错的根源,以及对应的解决方案:
报错原因解析
你遇到的InvalidArgumentError核心问题是:TensorFlow无法确定数据集的batch维度大小,导致MirroredStrategy无法将batch正确分配到多个GPU上。具体有三个直接诱因:
- 预处理函数错误:你的
preprocess函数给每个单样本添加了额外的维度(np.expand_dims(x, axis=0)),这会让生成器输出的图像batch形状变成(batch_size, 1, 224, 224, 3),完全不符合模型输入的预期,也让TensorFlow无法识别正确的维度结构。 - tf.data.Dataset定义错误:
flow_from_directory返回的生成器每次输出的是**(图像batch, 标签batch)**的元组,但你在from_generator里只指定了单一图像的形状,没处理标签部分,TensorFlow无法识别完整的输出结构,自然无法确定batch维度。 - 变量作用域问题:你在
strategy.scope()里引用了train_generator和test_generator,但这两个变量只在嵌套函数内部定义,main函数的当前作用域里根本没有这些变量,运行到那一行会触发NameError。
分步解决方案
1. 修正预处理函数
去掉多余的np.expand_dims,因为ImageDataGenerator会自动将单个样本组合成batch,不需要手动增加维度:
def preprocess(x): x /= 255.0 x -= 0.5 x *= 2.0 return x
2. 正确创建tf.data.Dataset
要明确告诉TensorFlow生成器输出的是图像和标签的元组,并且指定它们的形状。同时要保存生成器实例,方便获取样本数量和类别数:
def main(args): # 保留原有callback定义 log_dir = args.log_dir#'logs/' checkpoint = ModelCheckpoint(args.model_save_dir + 'ep{epoch:03d}-val_loss{val_loss:.3f}-val_acc{val_acc:.3f}-val_top_k_categorical_accuracy{val_top_k_categorical_accuracy:.3f}.h5', monitor='val_acc', mode='max', verbose=1, save_weights_only=False, save_best_only=True, period=1) logging = TensorBoard(log_dir=args.model_save_dir, histogram_freq=0, write_graph=False, write_grads=False, write_images=False, update_freq='batch') terminate_on_nan = TerminateOnNaN() learn_rates = [0.05, 0.01, 0.005, 0.001, 0.0005, 0.0001] lr_scheduler = LearningRateScheduler(lambda epoch: learn_rates[epoch // 30]) # ---------------------- 修正训练数据集部分 ---------------------- train_datagen = ImageDataGenerator(preprocessing_function=preprocess, zoom_range=0.25, width_shift_range=0.05, height_shift_range=0.05, horizontal_flip=True) # 使用全局batch size(单GPU batch数 * GPU数量) train_generator = train_datagen.flow_from_directory( args.train_data_path, target_size=(224, 224), batch_size=batch_size_all) total_train_samples = train_generator.samples num_classes = train_generator.num_classes # 获取数据集类别数 # 包装生成器,确保输出是(图像, 标签)元组 def train_gen_wrapper(): for images, labels in train_generator: yield images, labels # 正确定义Dataset的输出类型和形状 train_dataset = tf.data.Dataset.from_generator( train_gen_wrapper, output_types=(tf.float32, tf.float32), output_shapes=(tf.TensorShape([None, 224, 224, 3]), tf.TensorShape([None, num_classes])) ) # ---------------------- 修正验证数据集部分 ---------------------- test_datagen = ImageDataGenerator(preprocessing_function=preprocess) test_generator = test_datagen.flow_from_directory( args.val_data_path, target_size=(224, 224), batch_size=batch_size_all) total_test_samples = test_generator.samples def test_gen_wrapper(): for images, labels in test_generator: yield images, labels test_dataset = tf.data.Dataset.from_generator( test_gen_wrapper, output_types=(tf.float32, tf.float32), output_shapes=(tf.TensorShape([None, 224, 224, 3]), tf.TensorShape([None, num_classes])) ) # ---------------------- 模型训练部分 ---------------------- optimizer = get_optimizer(args.optim_type, args.learning_rate) with strategy.scope(): model = get_model(args.model_type) model.compile( optimizer=optimizer, metrics=['accuracy', 'top_k_categorical_accuracy'], loss='categorical_crossentropy') model.summary() # 现在可以正确输出样本信息了 print('Train on {} samples, val on {} samples, with batch size {}.'.format( total_train_samples, total_test_samples, batch_size_all)) # 计算训练和验证步数(总样本数 // 全局batch size) steps_per_epoch = total_train_samples // batch_size_all validation_steps = total_test_samples // batch_size_all model.fit( train_dataset, batch_size=None, epochs=args.total_epoch, initial_epoch=args.init_epoch, validation_data=test_dataset, steps_per_epoch=steps_per_epoch, validation_steps=validation_steps, callbacks=[logging, checkpoint, lr_scheduler, terminate_on_nan]) model.save(log_dir + 'trained_final.h5')
3. 关键注意点
- 确保
get_model返回的模型输入形状是(224, 224, 3),和预处理后的图像维度匹配。 batch_size_all要和flow_from_directory的batch_size参数保持一致,这样生成器每次输出的batch大小才符合MirroredStrategy的要求。- 必须指定
steps_per_epoch和validation_steps:因为从生成器创建的tf.data.Dataset无法自动推断迭代次数,需要手动计算总样本数除以全局batch size。
这样修改后,TensorFlow就能明确识别数据集的结构,MirroredStrategy可以正常地将batch拆分到多个GPU上进行训练,解决你遇到的报错问题。
内容的提问来源于stack exchange,提问作者junjun
相关产品推荐
相关产品推荐

