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

Model.fit()与Model.predict()结果不一致及验证集标签异常问题

问题根源与解决方案

核心问题

准确率差异和标签结果不稳定的本质是两次遍历val_ds时样本顺序不匹配:

  • model.predict(val_ds)会遍历一遍验证集生成预测结果;之后再次遍历val_ds提取标签时,由于val_ds(SkipDataset)的上游数据集可能启用了无固定种子的shuffle,或未做缓存,导致两次遍历的样本顺序完全不同,预测结果和标签错位,计算出错误的低准确率。
  • 每次提取标签结果不同,直接验证了val_ds的迭代顺序是随机变化的。

具体解决方案

1. 缓存验证集,固定迭代顺序

在创建val_ds后添加缓存操作,让数据集只生成一次,后续迭代复用缓存内容,确保顺序一致:

# 内存缓存(适合小数据集)
val_ds = val_ds.cache()

# 磁盘缓存(适合大数据集)
val_ds = val_ds.cache("./validation_cache")

2. 一次性提取所有验证数据,避免重复遍历

不要分两次遍历val_ds,而是一次性获取所有样本和标签,再进行预测和准确率计算:

# 一次性提取验证集的所有图像和标签
val_images = np.concatenate([x for x, y in val_ds], axis=0)
val_labels = np.concatenate([y for x, y in val_ds], axis=0)

# 基于提取的图像进行预测
cnn1_pred = model.predict(val_images).argmax(axis=-1)

# 计算准确率(简化写法)
correct = np.sum(val_labels == cnn1_pred)
perf = round(correct / len(val_labels), 4)

3. 固定shuffle的随机种子(如果使用了shuffle)

如果val_ds的上游数据集启用了shuffle,必须指定固定的seed参数,强制每次迭代顺序一致:

# 示例:创建val_ds时固定shuffle种子
val_ds = original_dataset.skip(TRAIN_SIZE).shuffle(buffer_size=1000, seed=42).batch(BATCH_SIZE)

额外验证:确保模型处于推理模式

虽然model.predict默认会自动禁用Dropout等训练层,但可以显式指定training=False确保推理模式:

cnn1_pred = model(val_images, training=False).numpy().argmax(axis=-1)

内容的提问来源于stack exchange,提问作者Daniel_Fortesque

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 21:18:33