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

如何将花卉数据集拆分为80:10:10的训练、验证、测试集?

问题描述

我正在使用一个按类别文件夹组织的花卉数据集——每个类别对应独立的子文件夹,文件夹内存放该类别的所有花卉图片。当前已将其按80:20的比例拆分为训练集和验证集并完成了网络训练,现在希望调整为80%训练集、10%验证集、10%测试集的拆分比例,并用TensorFlow的model.evaluate()方法测试模型。

现有代码如下:

import pathlib
dataset_url = "https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz"
data_dir = tf.keras.utils.get_file(origin=dataset_url,
                                   fname='flower_photos',
                                   untar=True)
data_dir = pathlib.Path(data_dir)
# Loader params
batch_size = 32
img_height = 180
img_width = 180
# Training imgs
train_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="training",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)
# Validation imgs
val_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="validation",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

我曾尝试在创建训练/验证集前手动提取图片但未成功,想知道更简便的实现方法。

解决方案

可以通过两次分步拆分或者直接用tf.data.Dataset的拆分方法实现,以下是两种贴合你现有代码的简便方案:

方案1:基于image_dataset_from_directory分步拆分(推荐)

先从全量数据中拆分出90%的「训练+验证集」和10%的测试集,再从「训练+验证集」里拆分出8/9(对应全量的80%)作为训练集,1/9(对应全量的10%)作为验证集,保证比例精准:

import pathlib
import tensorflow as tf  # 补充原代码缺失的TensorFlow导入

dataset_url = "https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz"
data_dir = tf.keras.utils.get_file(origin=dataset_url,
                                   fname='flower_photos',
                                   untar=True)
data_dir = pathlib.Path(data_dir)

# Loader params
batch_size = 32
img_height = 180
img_width = 180
seed = 123

# 第一步:拆分出90%的train_val_ds和10%的test_ds
train_val_ds = tf.keras.utils.image_dataset_from_directory(
    data_dir,
    validation_split=0.1,  # 预留10%作为测试集
    subset="training",
    seed=seed,
    image_size=(img_height, img_width),
    batch_size=batch_size
)

test_ds = tf.keras.utils.image_dataset_from_directory(
    data_dir,
    validation_split=0.1,
    subset="validation",
    seed=seed,
    image_size=(img_height, img_width),
    batch_size=batch_size
)

# 第二步:从train_val_ds拆分出80%训练集和10%验证集
total_train_val = len(train_val_ds) * batch_size
train_size = int(total_train_val * (8/9))
# 打乱数据集保证类别分布均匀
train_val_ds = train_val_ds.shuffle(buffer_size=total_train_val, seed=seed)
train_ds = train_val_ds.take(train_size // batch_size)
val_ds = train_val_ds.skip(train_size // batch_size)

# 可选:优化数据集加载性能
AUTOTUNE = tf.data.AUTOTUNE
train_ds = train_ds.cache().shuffle(1000).prefetch(buffer_size=AUTOTUNE)
val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)
test_ds = test_ds.cache().prefetch(buffer_size=AUTOTUNE)

方案2:直接拆分全量数据集

如果不想多次调用加载函数,可以先加载全量数据集,再按比例拆分:

import pathlib
import tensorflow as tf

dataset_url = "https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz"
data_dir = tf.keras.utils.get_file(origin=dataset_url,
                                   fname='flower_photos',
                                   untar=True)
data_dir = pathlib.Path(data_dir)

# Loader params
batch_size = 32
img_height = 180
img_width = 180
seed = 123

# 加载全量数据集
full_ds = tf.keras.utils.image_dataset_from_directory(
    data_dir,
    seed=seed,
    image_size=(img_height, img_width),
    batch_size=batch_size
)

# 计算各数据集的样本量
total_samples = len(full_ds) * batch_size
train_size = int(0.8 * total_samples)
val_size = int(0.1 * total_samples)

# 打乱后按比例拆分
full_ds = full_ds.shuffle(buffer_size=total_samples, seed=seed)
train_ds = full_ds.take(train_size // batch_size)
val_ds = full_ds.skip(train_size // batch_size).take(val_size // batch_size)
test_ds = full_ds.skip(train_size // batch_size + val_size // batch_size)

# 优化数据集性能
AUTOTUNE = tf.data.AUTOTUNE
train_ds = train_ds.cache().shuffle(1000).prefetch(buffer_size=AUTOTUNE)
val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)
test_ds = test_ds.cache().prefetch(buffer_size=AUTOTUNE)

模型测试

拆分完成后,直接调用model.evaluate()即可完成测试:

test_loss, test_acc = model.evaluate(test_ds, verbose=2)
print(f"测试准确率: {test_acc}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 00:31:04