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

如何在TensorFlow中先打乱再拆分PrefetchDataset为训练集和测试集

解决方案

1. 先打乱数据集

你可以直接调用数据集的shuffle()方法完成打乱操作,关键是设置合适的buffer_size(建议设为数据集总样本数的1/10至全量,至少远大于batch size,保证打乱效果)。另外建议加载数据集时先不设置num_epochs,拆分后再给训练集添加重复逻辑,避免重复处理整个数据集。

代码示例:

import tensorflow as tf

# 加载CSV数据集,暂不指定num_epochs
data = tf.data.experimental.make_csv_dataset("flight_2018.csv",
                                             batch_size=1000,
                                             label_name="Cancelled",
                                             num_parallel_reads=2)

# 打乱数据集,buffer_size根据实际样本量调整(这里示例设为100000)
shuffled_data = data.shuffle(buffer_size=100000, reshuffle_each_iteration=True)

reshuffle_each_iteration=True会让每个epoch都重新打乱数据,更适合训练场景。

2. 拆分训练集与测试集

通过take()(取前N个批次)和skip()(跳过前N个批次)方法按比例拆分,比如常用的8:2划分:

步骤1:计算总批次数量

# 遍历数据集统计总批次(batch_size=1000,总样本数=总批次×1000)
total_batches = sum(1 for _ in shuffled_data)
# 按8:2比例分配训练/测试批次
train_batches = int(total_batches * 0.8)

步骤2:执行拆分并设置epochs

# 拆分训练集,同时设置训练需要的20个epochs
train_dataset = shuffled_data.take(train_batches).repeat(20)
# 拆分测试集,测试集一般无需重复
test_dataset = shuffled_data.skip(train_batches)

完整代码示例

import tensorflow as tf

# 加载数据集
data = tf.data.experimental.make_csv_dataset("flight_2018.csv",
                                             batch_size=1000,
                                             label_name="Cancelled",
                                             num_parallel_reads=2)

# 打乱数据
shuffled_data = data.shuffle(buffer_size=100000, reshuffle_each_iteration=True)

# 计算批次并拆分
total_batches = sum(1 for _ in shuffled_data)
train_batches = int(total_batches * 0.8)

train_dataset = shuffled_data.take(train_batches).repeat(20)
test_dataset = shuffled_data.skip(train_batches)

注意事项

  • buffer_size不能太小,否则打乱效果差;也无需过大,避免占用过多内存,针对280MB的CSV,10万-20万的buffer_size通常足够。
  • 必须先打乱再拆分,否则会出现训练/测试集的数据分布偏差(比如前80%是早期数据,后20%是后期数据)。
  • 如果统计总批次耗时,可直接估算总样本数后计算批次,比如已知总样本约100万,则训练批次为int(1000000/1000*0.8)=800。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 13:55:20