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

如何在Keras数据生成器中划分训练-验证-测试数据集?

Keras数据生成器的训练-验证-测试数据集划分方案

首先纠正你当前代码的核心问题:你直接把测试集文件夹用作了验证集,这会让测试集数据参与模型训练过程中的性能监控与调优,最终用它做预测得到的结果无法真实反映模型的泛化能力。

正确的数据集划分逻辑是:

  • 训练集:用于更新模型参数
  • 验证集:训练过程中监控性能、调整超参数(如早停、学习率)
  • 测试集:仅在模型最终训练完成后,评估真实泛化能力,全程不参与训练和调优

两种可行实现方案

方案1:从训练文件夹拆分验证集(推荐)

利用ImageDataGenerator的validation_split参数,直接从训练数据中拆分出验证子集,无需额外创建验证文件夹:

# 训练集生成器(可添加数据增强,验证/测试集仅做归一化)
train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=20,
    width_shift_range=0.2,
    height_shift_range=0.2,
    horizontal_flip=True,
    validation_split=0.2  # 划分20%训练数据为验证集
)

# 验证、测试集仅做归一化处理
valid_test_datagen = ImageDataGenerator(rescale=1./255)

# 加载训练子集
train_generator = train_datagen.flow_from_directory(
    train_dir,
    target_size=(224, 224),
    class_mode='categorical',
    subset='training'
)

# 加载从训练集拆分出的验证子集
valid_generator = train_datagen.flow_from_directory(
    train_dir,
    target_size=(224, 224),
    class_mode='categorical',
    subset='validation'
)

# 加载独立的测试集
test_generator = valid_test_datagen.flow_from_directory(
    test_dir,
    target_size=(224, 224),
    class_mode='categorical',
    shuffle=False  # 关闭洗牌,保证预测结果与标签顺序对应
)

# 模型训练
history = model.fit(
    train_generator,
    epochs=10,
    validation_data=valid_generator,
    verbose=1
)

# 最终预测使用测试集生成器
pred = model.predict(test_generator)

方案2:提前划分三个独立文件夹

如果已经手动将数据按比例分成训练、验证、测试三个独立文件夹,直接分别加载即可:

train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=20,
    horizontal_flip=True
)
valid_datagen = ImageDataGenerator(rescale=1./255)
test_datagen = ImageDataGenerator(rescale=1./255)

train_generator = train_datagen.flow_from_directory(
    train_dir,
    target_size=(224, 224),
    class_mode='categorical'
)

valid_generator = valid_datagen.flow_from_directory(
    val_dir,  # 单独的验证文件夹
    target_size=(224, 224),
    class_mode='categorical'
)

test_generator = test_datagen.flow_from_directory(
    test_dir,
    target_size=(224, 224),
    class_mode='categorical',
    shuffle=False
)

history = model.fit(
    train_generator,
    epochs=10,
    validation_data=valid_generator,
    verbose=1
)

# 用测试集生成器做最终预测
pred = model.predict(test_generator)

关键注意事项

  • 测试集绝对不能用作训练时的validation_data,否则模型会在训练中“见过”测试数据,导致泛化能力评估失真。
  • 预测测试集时建议设置shuffle=False,保证预测结果顺序与test_generator.classes一致,方便后续计算准确率、混淆矩阵等指标。
  • 验证集和测试集不要添加数据增强,仅做与训练集一致的归一化处理即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 05:15:32