如何将TensorFlow的PrefetchDataset拆分为训练集与测试集?
拆分PrefetchDataset为训练集和测试集的方法
你用tf.data.experimental.make_csv_dataset得到的data属于tf.data.Dataset子类,以下是两种实用的拆分方法:
方法一:按比例拆分(推荐)
无需提前统计样本总数,直接按比例划分训练/测试集(示例为8:2):
调整数据集配置
先将num_epochs设为None(让数据集可重复迭代),同时缓存数据集避免重复读取CSV,提升效率:import tensorflow as tf a = [0, 2, 3, 4, 5, 19, 23, 32, 39, 40, 42, 50, 51, 53, 56, 65, 66, 67, 68, 69] data = tf.data.experimental.make_csv_dataset("./raw/*", batch_size=2000, select_columns=a, label_name="Cancelled", num_epochs=None, num_parallel_reads=2) data = data.cache()执行拆分
先获取总批次数量,再按比例截取训练集和测试集,最后设置迭代epoch数:total_batches = len(list(data)) train_batches = int(0.8 * total_batches) test_batches = int(0.2 * total_batches) train_dataset = data.take(train_batches).repeat(30) test_dataset = data.skip(train_batches).take(test_batches).repeat(30)
方法二:按样本数拆分
如果已经统计出所有CSV的总样本数(比如100000条),可以按具体样本量拆分:
打散批次为单个样本
先把批次数据拆成单个样本,方便按样本数截取:data_single = data.unbatch()拆分并重新批次化
按预设样本数拆分后,重新组合为原批次大小,再设置epoch:# 假设训练集80000条,测试集20000条 train_dataset = data_single.take(80000).batch(2000).repeat(30) test_dataset = data_single.skip(80000).take(20000).batch(2000).repeat(30)
注意事项
- 拆分后的
train_dataset和test_dataset可以直接传入model.fit()等训练接口使用。 - 若数据集规模极大,
len(list(data))会遍历一次数据集,耗时较长,建议提前统计总样本数或使用tf.data.experimental.cardinality()尝试获取数据集大小(无限数据集会返回tf.data.UNKNOWN_CARDINALITY)。
内容的提问来源于stack exchange,提问作者Shawn Brar
相关产品推荐
相关产品推荐

