TensorFlow图像分类缓存警告及数据集未充分使用问题求助
TensorFlow图像分类数据集问题解决方向
1. 数据集划分参数错误
你当前的训练集和验证集比例完全不符合预期,核心是对validation_split参数的含义理解有误:
validation_split指定的是验证集占总数据的比例,而非训练集的比例- 你给
train_ds设置validation_split=0.8,意味着从总数据中划分80%作为验证集,剩下20%才是训练集;而val_ds设置validation_split=0.2,又从总数据中单独划分20%作为验证集。这就导致训练集仅占总数据的20%,且训练集和验证集大概率存在数据重叠。
修正后的数据集生成代码:
#Set up information on the data batch_size = 32 img_height = 100 img_width = 100 #Generate training dataset train_ds = tf.keras.utils.image_dataset_from_directory( Directory, validation_split=0.2, # 验证集占20%,训练集自动分配剩余80% subset="training", seed=123, image_size=(img_height, img_width), batch_size=batch_size) #Generate val dataset val_ds = tf.keras.utils.image_dataset_from_directory( Directory, validation_split=0.2, # 保持和训练集一致的划分比例 subset="validation", seed=123, image_size=(img_height, img_width), batch_size=batch_size)
修正后训练集数据量应为2080581 * 0.8 ≈ 1664465,验证集为2080581 * 0.2 ≈ 416116,和输出的验证集数量匹配,训练集数量也会恢复正常。
2. 缓存警告处理
终端的警告核心是数据集缓存的执行顺序错误,导致迭代器未完全读取缓存内容就被截断。
如果你的后续代码中对数据集做了类似这样的操作:
train_ds = train_ds.cache().take(k).repeat()
就会触发该警告,正确的顺序应该是先执行take等可能截断数据集的操作,再做缓存:
train_ds = train_ds.take(k).cache().repeat()
如果是直接使用官方教程的流水线优化代码(如cache() + prefetch())但未做take操作,那可能是训练时的steps_per_epoch参数设置错误,或者epoch设置导致数据集未被完全遍历。此时可以:
- 移除手动指定的
steps_per_epoch,让TensorFlow自动计算 - 确认数据集batch数量与epoch的乘积符合训练逻辑
内容的提问来源于stack exchange,提问作者j.t.2.4.6
相关产品推荐
相关产品推荐

