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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 19:55:20