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

如何在TensorFlow中使用自有数据划分图像训练集与测试集

数据集划分方案:从train集拆分20%作为测试集

嘿,我来帮你理清楚这个数据集划分的问题!你完全不用在加载成tf.data.Dataset之前做划分——其实有两种靠谱的方式,我结合你的现有代码给你一步步讲:

方法一:用sklearn分层拆分(推荐,保证类别分布一致)

这种方法适合大多数场景,尤其是你的数据集存在类别不平衡的情况,能确保训练集和测试集的类别比例和原数据完全一致。

步骤1:把文件路径Dataset转成可迭代的列表

首先,我们需要把list_ds里的文件路径提取出来,因为train_test_split需要处理普通的序列数据:

from sklearn.model_selection import train_test_split
import numpy as np

# 提取所有文件路径到列表
file_paths = list(list_ds.as_numpy_iterator())

步骤2:分层拆分训练/测试路径

这里关键是用stratify参数,基于每个文件的类别来分层,避免某类样本在测试集里占比失衡:

# 从路径中提取每个文件对应的类别(匹配CLASS_NAMES的索引)
labels = []
for path in file_paths:
    # 解析路径最后第二个部分(类别文件夹名称)
    class_name = path.decode().split('/')[-2]
    label_idx = np.argmax(CLASS_NAMES == class_name)
    labels.append(label_idx)

# 拆分:测试集占20%,固定random_state保证结果可复现
train_paths, test_paths = train_test_split(
    file_paths,
    test_size=0.2,
    random_state=42,
    stratify=labels
)

步骤3:把拆分后的路径转回Dataset并处理图像/标签

先写一个解析函数,用来加载图像并生成对应的标签:

def parse_image(file_path):
    # 提取类别标签
    parts = tf.strings.split(file_path, '/')
    label = parts[-2] == CLASS_NAMES
    label = tf.argmax(label)  # 转成整数索引标签
    
    # 加载并预处理图像
    img = tf.io.read_file(file_path)
    img = tf.image.decode_png(img, channels=3)  # 如果是jpg格式,换成decode_jpeg
    img = tf.image.convert_image_dtype(img, tf.float32)  # 归一化到0-1区间
    img = tf.image.resize(img, [224, 224])  # 根据你的模型需求调整尺寸
    
    return img, label

然后把训练/测试路径映射成带标签的数据集,并添加优化操作:

BATCH_SIZE = 32

# 映射解析函数,并行加速
train_ds = tf.data.Dataset.from_tensor_slices(train_paths).map(
    parse_image, num_parallel_calls=tf.data.AUTOTUNE
)
test_ds = tf.data.Dataset.from_tensor_slices(test_paths).map(
    parse_image, num_parallel_calls=tf.data.AUTOTUNE
)

# 打乱、分批、预加载,提升训练效率
train_ds = train_ds.shuffle(buffer_size=len(train_paths)).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
test_ds = test_ds.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)

方法二:直接对tf.data.Dataset拆分(适合超大数据集)

如果你的数据集大到把所有路径加载到内存会卡顿,那可以直接对tf.data.Dataset进行拆分,不用依赖sklearn:

# 先打乱数据集,固定seed保证划分结果可复现
shuffled_ds = list_ds.shuffle(buffer_size=image_count, seed=42, reshuffle_each_iteration=False)

# 计算测试集大小(20%)
test_size = int(image_count * 0.2)

# 拆分数据集
test_ds = shuffled_ds.take(test_size)
train_ds = shuffled_ds.skip(test_size)

# 后续同样应用parse_image函数和batch/prefetch操作,和方法一一致
train_ds = train_ds.map(parse_image, num_parallel_calls=tf.data.AUTOTUNE).shuffle(1000).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
test_ds = test_ds.map(parse_image, num_parallel_calls=tf.data.AUTOTUNE).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)

两种方法的对比

  • 方法一:优势是能严格保证训练/测试集的类别分布一致,适合类别不平衡的数据集;缺点是需要把所有路径加载到内存,超大数据集可能吃不消。
  • 方法二:优势是内存友好,不用加载所有路径;缺点是如果原数据集没有充分打乱,可能出现测试集类别分布不均的情况(如果要分层,需要额外按类别分组后再拆分,步骤会更繁琐)。

根据你的情况,我更推荐方法一,毕竟能避免类别分布偏差,让模型评估更准确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 16:17:36