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

如何在MirroredStrategy下结合ImageDataGenerator.flow_from_directory使用model.fit()?

解决MirroredStrategy配合ImageDataGenerator转tf.data.Dataset的报错问题

我帮你分析下报错的根源,以及对应的解决方案:

报错原因解析

你遇到的InvalidArgumentError核心问题是:TensorFlow无法确定数据集的batch维度大小,导致MirroredStrategy无法将batch正确分配到多个GPU上。具体有三个直接诱因:

  1. 预处理函数错误:你的preprocess函数给每个单样本添加了额外的维度(np.expand_dims(x, axis=0)),这会让生成器输出的图像batch形状变成(batch_size, 1, 224, 224, 3),完全不符合模型输入的预期,也让TensorFlow无法识别正确的维度结构。
  2. tf.data.Dataset定义错误:flow_from_directory返回的生成器每次输出的是**(图像batch, 标签batch)**的元组,但你在from_generator里只指定了单一图像的形状,没处理标签部分,TensorFlow无法识别完整的输出结构,自然无法确定batch维度。
  3. 变量作用域问题:你在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 00:53:12