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

如何将TensorFlow的PrefetchDataset拆分为训练集与测试集?

拆分PrefetchDataset为训练集和测试集的方法

你用tf.data.experimental.make_csv_dataset得到的data属于tf.data.Dataset子类,以下是两种实用的拆分方法:

方法一:按比例拆分(推荐)

无需提前统计样本总数,直接按比例划分训练/测试集(示例为8:2):

  1. 调整数据集配置
    先将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()
    
  2. 执行拆分
    先获取总批次数量,再按比例截取训练集和测试集,最后设置迭代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条),可以按具体样本量拆分:

  1. 打散批次为单个样本
    先把批次数据拆成单个样本,方便按样本数截取:

    data_single = data.unbatch()
    
  2. 拆分并重新批次化
    按预设样本数拆分后,重新组合为原批次大小,再设置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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 09:50:31