分类神经网络中何时使用accuracy字符串或tf.keras.metrics.Accuracy()
问题原因
你遇到的准确率差异问题,核心是"accuracy"字符串和tf.keras.metrics.Accuracy()实例并非等价实现:
- 当你传入字符串
"accuracy"时,Keras会根据你当前使用的损失函数自动匹配对应的准确率实现:你的场景是二分类、损失为binary_crossentropy,因此Keras会自动调用tf.keras.metrics.BinaryAccuracy(),这个类会自动先把sigmoid输出的0~1概率按照默认阈值0.5转为0/1的类别预测,再和真实标签对比计算准确率,所以得到的结果是正常的1.0。 - 当你手动传入
tf.keras.metrics.Accuracy()时,这个类不会做自动的概率转类别操作,会直接拿你模型输出的0~1浮点概率值和真实的0/1整数标签做相等判断,浮点数和整数几乎不可能完全相等,因此计算出来的准确率永远是0。
如果你想用类实例实现和字符串"accuracy"完全一致的效果,修正写法如下:
model.compile(loss=tf.keras.losses.binary_crossentropy, optimizer=tf.keras.optimizers.Adam(), metrics=[tf.keras.metrics.BinaryAccuracy()])
两种写法的使用场景
字符串形式适用场景
常规的分类、回归任务,不需要自定义指标参数的时候,直接传字符串即可,Keras的自动匹配逻辑可以适配绝大多数标准场景,写法更简洁不易出错:
- 二分类+binary_crossentropy损失:自动匹配BinaryAccuracy
- 多分类+categorical_crossentropy损失:自动匹配CategoricalAccuracy
- 多分类+sparse_categorical_crossentropy损失:自动匹配SparseCategoricalAccuracy
完整类实例适用场景
- 需要自定义指标的超参数时:比如你想把二分类准确率的判断阈值从默认的0.5改成0.3,就需要手动传实例
tf.keras.metrics.BinaryAccuracy(threshold=0.3) - 使用Keras内置的非默认匹配的指标时:比如你要在二分类任务里计算召回率、精确率、AUC等,就需要手动传入对应类的实例
tf.keras.metrics.Recall()、tf.keras.metrics.AUC() - 使用自定义指标的场景:自己实现的指标类,需要实例化后传入metrics参数
内容的提问来源于stack exchange,提问作者guilerme_coppola84
相关产品推荐
相关产品推荐

