如何在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
相关产品推荐
相关产品推荐

