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

TensorFlow训练输入数据耗尽问题求助及代码排查

问题原因与解决方法

核心问题分析

警告提示训练时输入数据耗尽,本质是数据集生成器的迭代次数超过了实际可提供的批次数量,同时代码存在几处细节错误:

  1. steps_per_epoch手动计算值可能和实际训练集样本数不匹配:你设置STEPS_PER_EPOCH = 480//BATCH_SIZE,但如果train_test_split划分后的训练集样本数不是刚好480,就会导致生成器提前耗尽数据。
  2. fit_generator和evaluate_generator已被TensorFlow/Keras弃用,旧方法对生成器迭代的处理逻辑存在局限。
  3. 测试阶段evaluate_generator的steps=100设置错误:测试集只有120个样本,BATCH_SIZE=2,最多只能提供60个批次,设置100会直接耗尽测试数据。
  4. 模型结构冗余:最后一层Dense(5, softmax)之后额外添加的Flatten()是无效操作,会破坏输出层结构。

具体修复步骤

1. 移除模型冗余层

删除输出层后的Flatten(),保证模型结构正确:

model = Sequential()
model.add(Conv2D(NUM_FILTERS, (FILTER_SIZE, FILTER_SIZE), input_shape = (INPUT_SIZE, INPUT_SIZE, 3), activation = 'relu'))
model.add(MaxPooling2D(pool_size = (MAXPOOL_SIZE, MAXPOOL_SIZE)))
model.add(Conv2D(NUM_FILTERS, (FILTER_SIZE, FILTER_SIZE), activation = 'relu'))
model.add(MaxPooling2D(pool_size = (MAXPOOL_SIZE, MAXPOOL_SIZE)))
model.add(Flatten())
model.add(Dense(units = 128, activation = 'relu'))
model.add(Dropout(0.5))
model.add(Dense(units = 5, activation = 'softmax'))  # 移除后续冗余的Flatten()
model.compile(optimizer = 'adam', loss = 'SparseCategoricalCrossentropy', metrics = ['accuracy'])

2. 动态获取实际批次数量

不要手动计算批次,直接从生成器的属性中获取,确保和真实样本数完全匹配:

# 生成数据集后,获取训练/测试集的实际批次
STEPS_PER_EPOCH = training_set.samples // training_set.batch_size
TEST_STEPS = test_set.samples // test_set.batch_size

3. 使用官方推荐的fit和evaluate方法

替换已弃用的fit_generator和evaluate_generator,新方法会自动处理生成器的循环迭代:

# 训练模型
model.fit(training_set, 
          steps_per_epoch = STEPS_PER_EPOCH, 
          epochs = EPOCHS, 
          verbose=1)

# 评估模型
score = model.evaluate(test_set, steps=TEST_STEPS)

4. 确保生成器的迭代稳定性

在生成数据集时开启shuffle=True(训练集默认开启,测试集建议关闭),保证训练时数据的随机性和迭代的连续性:

training_set = training_data_generator.flow_from_directory(src+'Train/',
                                                target_size = (INPUT_SIZE, INPUT_SIZE),
                                                batch_size = BATCH_SIZE,
                                                class_mode = 'sparse',
                                                shuffle=True)

test_set = testing_data_generator.flow_from_directory(src+'Test/',
                                             target_size = (INPUT_SIZE, INPUT_SIZE),
                                             batch_size = BATCH_SIZE,
                                             class_mode='sparse',
                                             shuffle=False)

完整修复后的代码

import os
import random
import warnings
warnings.filterwarnings("ignore")
from ut2 import train_test_split

src = 'Dataset/corrosion/'

# 创建训练/测试文件夹(如果不存在)
if not os.path.isdir(src+'train/'):
    train_test_split(src)

from keras.models import Sequential
from keras.layers import Conv2D, MaxPooling2D
from keras.layers import Dropout, Flatten, Dense
from keras.preprocessing.image import ImageDataGenerator

# 定义超参数
FILTER_SIZE = 3
NUM_FILTERS = 32
INPUT_SIZE  = 200
MAXPOOL_SIZE = 2
BATCH_SIZE = 2
EPOCHS = 50

# 构建模型
model = Sequential()
model.add(Conv2D(NUM_FILTERS, (FILTER_SIZE, FILTER_SIZE), input_shape = (INPUT_SIZE, INPUT_SIZE, 3), activation = 'relu'))
model.add(MaxPooling2D(pool_size = (MAXPOOL_SIZE, MAXPOOL_SIZE)))
model.add(Conv2D(NUM_FILTERS, (FILTER_SIZE, FILTER_SIZE), activation = 'relu'))
model.add(MaxPooling2D(pool_size = (MAXPOOL_SIZE, MAXPOOL_SIZE)))
model.add(Flatten())
model.add(Dense(units = 128, activation = 'relu'))
model.add(Dropout(0.5))
model.add(Dense(units = 5, activation = 'softmax'))
model.compile(optimizer = 'adam', loss = 'SparseCategoricalCrossentropy', metrics = ['accuracy'])

# 数据生成器
training_data_generator = ImageDataGenerator(rescale = 1./255)
testing_data_generator = ImageDataGenerator(rescale = 1./255)

training_set = training_data_generator.flow_from_directory(src+'Train/',
                                                target_size = (INPUT_SIZE, INPUT_SIZE),
                                                batch_size = BATCH_SIZE,
                                                class_mode = 'sparse',
                                                shuffle=True)

test_set = testing_data_generator.flow_from_directory(src+'Test/',
                                             target_size = (INPUT_SIZE, INPUT_SIZE),
                                             batch_size = BATCH_SIZE,
                                             class_mode='sparse',
                                             shuffle=False)

# 获取实际批次数量
STEPS_PER_EPOCH = training_set.samples // training_set.batch_size
TEST_STEPS = test_set.samples // test_set.batch_size

# 训练模型
model.fit(training_set, 
          steps_per_epoch = STEPS_PER_EPOCH, 
          epochs = EPOCHS, 
          verbose=1)

# 评估模型
score = model.evaluate(test_set, steps=TEST_STEPS)

for idx, metric in enumerate(model.metrics_names):
    print("{}: {}".format(metric, score[idx]))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 03:10:54