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

TensorFlow中model.predict准确率与model.fit的val_accuracy不匹配问题

为什么model.predict计算的验证集准确率和model.fit输出的val_accuracy不一致?

我使用tf.data.Dataset处理图像数据,训练时model.fit输出的最终val_accuracy为0.8580,但手动用model.predict计算的验证集准确率仅为0.798,预期两者结果一致却出现差异。已确认验证集设置了seed、真实标签获取方式正确、模型softmax输出符合概率分布,想知道忽略了什么细节。

相关代码与输出

验证集构建

val_ds = tf.keras.utils.image_dataset_from_directory(
    'my_path',
    validation_split=0.2,
    subset="validation",
    seed=38,
    image_size=(SIZE,SIZE),
)

数据集预处理

train_ds = train_ds.prefetch(buffer_size=AUTOTUNE)
val_ds = val_ds.prefetch(buffer_size=AUTOTUNE)

获取验证集真实标签

true_categories = tf.concat([y for x, y in val_ds], axis=0)

模型结构

inputs = tf.keras.Input(shape=(SIZE, SIZE, 3))
# ... 其他层
outputs = tf.keras.layers.Dense(len(CLASS_NAMES), activation=tf.keras.activations.softmax)(intermediate)
model = tf.keras.Model(inputs, outputs)

模型编译

model.compile(
  optimizer='adam', 
  loss=tf.keras.losses.SparseCategoricalCrossentropy(), 
  metrics=['accuracy'])

训练代码

history = model.fit(
  train_ds,
  validation_data=val_ds,
  epochs=10, 
  class_weight=class_weights) # 因类别不平衡使用类别权重

训练最后一轮输出

Epoch 10: val_accuracy did not improve from 0.92291
176/176 [==============================] - 191s 1s/step - loss: 0.9876 - accuracy: 0.7318 - val_loss: 0.4650 - val_accuracy: 0.8580

手动计算准确率代码

predictions = model.predict(val_ds, verbose=2 ) 
flattened_predictions = predictions.argmax(axis=1)
accuracy = metrics.accuracy_score(true_categories, flattened_predictions)
print ("Accuracy = ", accuracy)

手动计算结果

Accuracy = 0.7980014275517487

问题原因

核心问题是验证集的shuffle配置导致两次遍历顺序不一致:

  • image_dataset_from_directory默认开启shuffle=True,且默认参数reshuffle_each_iteration=True,这意味着每次遍历数据集时都会重新打乱数据顺序。
  • 第一次遍历val_ds获取true_categories时,数据是一种顺序;第二次调用model.predict(val_ds)时,数据集被重新打乱,预测结果的顺序和之前获取的标签顺序不匹配,导致计算出的准确率错误。

解决方案

有两种方式确保两次遍历的顺序一致:

  1. 关闭验证集的shuffle(推荐,验证集通常不需要打乱):
val_ds = tf.keras.utils.image_dataset_from_directory(
    'my_path',
    validation_split=0.2,
    subset="validation",
    seed=38,
    image_size=(SIZE,SIZE),
    shuffle=False  # 关闭打乱,确保每次遍历顺序完全一致
)
val_ds = val_ds.prefetch(buffer_size=AUTOTUNE)
  1. 保留shuffle但固定打乱顺序:
    如果确实需要对验证集进行shuffle,可以设置reshuffle_each_iteration=False并缓存数据集,确保每次迭代使用相同的打乱顺序:
val_ds = tf.keras.utils.image_dataset_from_directory(
    'my_path',
    validation_split=0.2,
    subset="validation",
    seed=38,
    image_size=(SIZE,SIZE),
    shuffle=True,
    reshuffle_each_iteration=False  # 禁用每次迭代重新打乱
).cache().prefetch(buffer_size=AUTOTUNE)

额外验证

修改后可以检查true_categories和flattened_predictions的对应关系,比如打印前10个标签和预测值,确认顺序匹配,再重新计算准确率即可和model.fit的val_accuracy一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 04:09:16