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

Keras数据增强时Steps per epoch异常及相关技术问题咨询

图像增强与Keras训练问题解答

问题背景

我在使用图像增强时,原本将steps_per_epoch设置为原始的2-3倍没有问题,但这次出现了训练数据耗尽的报错。另外,修改Keras中的batchsize后,steps per epoch始终显示为193,没有变化。

相关配置

  • 训练数据量:6149
  • 验证数据量:1020
  • 尝试过的batchsize:32、64

数据增强与预处理代码

from tensorflow.keras.preprocessing.image import ImageDataGenerator
from tensorflow.keras.applications.resnet import preprocess_input
datagen = ImageDataGenerator(rotation_range=90,
                             horizontal_flip=True,
                             vertical_flip=True,
                             # featurewise_center=True,
                             zoom_range=0.3,
                             shear_range=0.5,
                             preprocessing_function=preprocess_input)
train_generator = datagen.flow_from_directory('data/train', seed=42)
validation_generator = datagen.flow_from_directory('data/validation', seed=42)

输出:

Found 6149 images belonging to 102 classes.
Found 1020 images belonging to 102 classes.

模型训练代码

batch_size = 64  # or 32
model.fit(train_generator, 
          epochs=100,
          batch_size=batch_size,
          # steps_per_epoch=3*training_size//batch_size, - gives error
          validation_data=validation_generator,
          callbacks=[early_stopping])

batchsize设为32或64时的输出示例:

Epoch 1/100
193/193 [==============================] - 3284s 17s/step - loss: 0.3660 - categorical_accuracy: 0.5357 - val_loss: 0.4022 - val_categorical_accuracy: 0.5176

指定steps_per_epoch时的报错:

Your input ran out of data; interrupting training. Make sure that your dataset or generator can generate at least steps_per_epoch * epochs batches (in this case, 57600 batches). You may need to use the repeat() function when building your dataset


问题解答

1. 未指定steps_per_epoch时,为何修改batchsize后steps per epoch始终为193?

你在model.fit()中设置的batch_size参数无效——flow_from_directory()生成的迭代器会使用自身默认的batch_size(默认值为32),而非model.fit()中的设置。你的训练数据共6149张,6149 ÷ 32 ≈ 192.16,向上取整后就是193个step,因此无论你在model.fit()中写32还是64,实际用的都是默认的32 batch size,steps数自然不变。

要修改batch size,必须在flow_from_directory()中显式指定:

train_generator = datagen.flow_from_directory('data/train', seed=42, batch_size=64)

2. 指定steps_per_epoch后为何出现训练数据不足的报错?

默认情况下,ImageDataGenerator生成的迭代器在一个epoch内只会遍历一次原始训练数据,增强样本是基于这批原始数据实时生成的,但不会重复使用原始数据。

当你设置steps_per_epoch=3*training_size//batch_size时,相当于要求每个epoch生成3倍于原始数据量的batch。但原始数据只能提供training_size//batch_size个batch,超过这个数量后,迭代器没有更多原始数据可生成增强样本,因此抛出数据耗尽的错误。

解决方法:

  • 给生成器添加.repeat(),让它无限循环使用原始数据生成新的增强样本:
    train_generator = datagen.flow_from_directory('data/train', seed=42, batch_size=32).repeat()
    
  • 或者保持steps_per_epoch等于原始数据的batch数(6149//batch_size),通过增加epoch数来获取更多增强样本的训练机会。

3. 我的生成器代码是否存在问题?

有几处需要调整:

  1. 未显式指定batch_size:如问题1所述,需在flow_from_directory()中指定,否则model.fit()的batch_size设置无效。
  2. 验证集使用数据增强:验证集用于评估模型真实表现,无需做数据增强,否则会导致验证结果不准确。应单独创建仅含预处理的验证生成器:
    val_datagen = ImageDataGenerator(preprocessing_function=preprocess_input)
    validation_generator = val_datagen.flow_from_directory('data/validation', seed=42, batch_size=32)
    
  3. 未指定target_size:flow_from_directory()默认输出尺寸为(256,256),但ResNet系列模型的标准输入尺寸是(224,224),需显式指定以匹配模型要求:
    train_generator = datagen.flow_from_directory('data/train', seed=42, batch_size=32, target_size=(224,224))
    
  4. 训练与验证生成器使用相同seed:虽不会直接报错,但可能导致验证集增强模式与训练集过于相似,建议给验证生成器设置不同seed。

4. 任务有102个类别,batchsize是否需要大于类别数?且该任务一直存在过拟合现象。

  • batchsize不需要大于类别数:没有规则要求batch size必须超过类别数,尤其是类别数较多的场景(如102类),batch size通常远小于类别数,只要保证每个batch能覆盖部分类别即可。若batch size过小导致类别覆盖不足,可适当调大,但无需超过102。
  • 过拟合解决建议:
    • 强化数据增强:添加随机裁剪、亮度调整等更多增强方式,或使用tf.keras.layers.RandomAugmentation工具。
    • 预训练模型冻结:基于ResNet预训练权重初始化,冻结前几层仅训练顶层,减少可训练参数。
    • 添加正则化:在全连接层或卷积层加入kernel_regularizer=tf.keras.regularizers.l2(0.01),或添加Dropout层。
    • 简化模型结构:减少全连接层数量或降低神经元数量。
    • 数据扩充:收集更多真实数据,或使用MixUp/CutMix等数据合成方法。
    • 优化早停回调:确保early_stopping监控验证集损失,及时停止训练。
    • 调整学习率:使用学习率衰减或AdamW等自适应优化器,避免模型过快拟合。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 05:50:41