使用tf.data的take、skip拆分数据集后验证集测试集为空是什么原因
问题核心原因
- 数据集拆分时计数单位不匹配
你在测试训练集时输出的Image shape第一维为32,说明你在执行拆分逻辑之前,已经对原始数据集ds调用了batch(32)方法,此时ds的每个元素对应一个包含32个样本的批次,而非单个图片样本。
你定义的DATASET_SIZE = 2000是单张样本的总数量,但tf.data.Dataset的take()、skip()方法的计数单位是数据集本身的元素个数,此时你的ds总元素数仅为2000//32 + (1 if 2000%32 else 0) = 63个(即总共有63个批次)。
执行ds.take(1400)时会直接把全部63个批次都划入训练集,后续的skip(1400)操作没有剩余元素可取,因此val_ds和test_ds均为空数据集。 - 若你已经在拆分前完成了shuffle操作,未固定随机种子也可能导致不同次迭代时数据集元素顺序不一致,出现部分拆分结果为空的情况,但该问题出现概率远低于单位不匹配问题。
正确实现逻辑
数据集拆分需要在batch操作之前完成,以单样本为单位执行拆分后,再对各个子集单独做batch、预处理、预加载等操作,示例代码如下:
import tensorflow as tf DATASET_SIZE = 2000 BATCH_SIZE = 32 train_size = int(0.7 * DATASET_SIZE) # 1400 val_size = int(0.15 * DATASET_SIZE) # 300 test_size = int(0.15 * DATASET_SIZE) # 300 # 加载原始单样本数据集,禁用自动batch # 若你使用image_dataset_from_directory加载,需设置batch_size=None ds = tf.keras.utils.image_dataset_from_directory( "你的数据集路径", batch_size=None, # 核心参数,返回单样本而非批次 image_size=(400, 400) ) # 打乱数据集,固定seed保证拆分结果可复现 ds = ds.shuffle(buffer_size=DATASET_SIZE, seed=42) # 先拆分单样本数据集 train_ds = ds.take(train_size) val_ds = ds.skip(train_size).take(val_size) test_ds = ds.skip(train_size + val_size).take(test_size) # 拆分完成后再分别执行batch和预加载 train_ds = train_ds.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE) val_ds = val_ds.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE) test_ds = test_ds.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
内容的提问来源于stack exchange,提问作者MAK
相关产品推荐
相关产品推荐

