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

多分类图像模型评估报错:Shapes (32,16)与(32,)不兼容

解决Keras多分类中Precision.update_state的形状不兼容错误

错误核心原因

ValueError: Shapes (32, 16) and (32,) are incompatible的本质是:

  • 模型输出yhat为one-hot编码格式(形状(32,16),对应16个分类),但测试集标签y是类别索引格式(形状`(32,)),两者数据格式不匹配。
  • 你误用了二分类专属的BinaryAccuracy指标,同时默认的Precision/Recall未配置多分类模式。

方案1:将标签转为one-hot编码匹配模型输出

把测试集标签转换为one-hot格式,同时替换为多分类适用的指标:

import tensorflow as tf
from tensorflow.keras.metrics import Precision, Recall, CategoricalAccuracy

# 原模型结构保持不变,此处省略

# 初始化多分类指标,指定average策略(按需选择macro/micro/weighted)
pre = Precision(average='macro')
re = Recall(average='macro')
ba = CategoricalAccuracy()

for batch in test.as_numpy_iterator():
    X, y = batch
    # 将类别索引转为one-hot编码,depth对应分类总数16
    y_one_hot = tf.one_hot(y, depth=16)
    yhat = model.predict(X, verbose=0)  # verbose=0关闭预测日志输出
    pre.update_state(y_one_hot, yhat)
    re.update_state(y_one_hot, yhat)
    ba.update_state(y_one_hot, yhat)

# 打印最终指标结果
print(f"Precision: {pre.result().numpy()}")
print(f"Recall: {re.result().numpy()}")
print(f"Accuracy: {ba.result().numpy()}")

方案2:将模型输出转为类别索引匹配标签

若不想修改标签格式,可将模型的概率输出转为类别索引,同时使用稀疏指标:

from tensorflow.keras.metrics import Precision, Recall, SparseCategoricalAccuracy

# 原模型结构保持不变,此处省略

# 初始化稀疏多分类指标(适配类别索引格式的标签)
pre = Precision(average='macro')
re = Recall(average='macro')
ba = SparseCategoricalAccuracy()

for batch in test.as_numpy_iterator():
    X, y = batch
    yhat = model.predict(X, verbose=0)
    # 将one-hot概率输出转为类别索引(取概率最大的类别)
    yhat_class = tf.argmax(yhat, axis=1)
    pre.update_state(y, yhat_class)
    re.update_state(y, yhat_class)
    ba.update_state(y, yhat)  # SparseCategoricalAccuracy可直接兼容两种格式

# 打印最终指标结果
print(f"Precision: {pre.result().numpy()}")
print(f"Recall: {re.result().numpy()}")
print(f"Accuracy: {ba.result().numpy()}")

关键注意事项

  • 多分类场景禁止使用BinaryAccuracy,该指标仅适用于二分类任务,多分类请选择:
    • CategoricalAccuracy:适配one-hot格式标签
    • SparseCategoricalAccuracy:适配类别索引格式标签
  • Precision/Recall在多分类时必须指定average参数:
    • 'macro':计算每类指标后取算术平均
    • 'micro':先统计全局TP/FP/FN再计算指标
    • 'weighted':按各类别样本数量加权平均

内容的提问来源于stack exchange,提问作者Likhith Chakravarthi K C

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 00:05:48