为何tf.keras.metrics.Accuracy()报错但metrics=['accuracy']可正常运行?
问题原因及解答
1. 报错根因
你遇到的形状不兼容报错,核心是Keras对字符串形式的指标和直接实例化的指标类处理逻辑不同:
- 当你传入字符串
"accuracy"作为metrics参数时,Keras会自动根据当前使用的损失函数类型匹配对应的指标实现:你代码中使用的是SparseCategoricalCrossentropy(稀疏分类交叉熵,对应标签为整数格式、非one-hot格式),Keras会自动调用tf.keras.metrics.SparseCategoricalAccuracy()类做计算,这个类会自动对模型输出的10维logits做argmax操作,转成和标签形状一致的预测类别整数,再做准确率计算,因此不会有形状问题。 - 当你直接实例化
tf.keras.metrics.Accuracy()传入时,这是通用准确率计算类,不会自动做logits转预测类别的处理,它要求输入的预测值y_pred和真实标签y_true形状完全一致。你模型输出的y_pred形状为(batch_size, 10)(每个样本对应10个类别的输出值),而真实标签y_true形状为(batch_size,)(每个样本对应1个整数类别),二者形状不匹配,因此抛出报错。
2. 二者功能是否一致
二者功能本质可以一致,但前提是你要手动匹配到正确的指标类:
你如果要手动实例化指标类实现和"accuracy"字符串完全相同的效果,只需要将代码中的指标替换为tf.keras.metrics.SparseCategoricalAccuracy()即可,修改后的compile代码如下:
model.compile( optimizer=tf.keras.optimizers.Adam(), loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=[tf.keras.metrics.SparseCategoricalAccuracy()] )
上述代码运行效果和传入"accuracy"字符串完全一致。
内容的提问来源于stack exchange,提问作者MachineLeon
相关产品推荐
相关产品推荐

