使用TensorFlow Metrics计算Precision始终为0,自行计算Accuracy正常求排查
这不是TensorFlow的Bug,是参数格式没匹配对
别慌,这绝对不是TensorFlow的问题,是你调用tf.metrics.precision时的参数格式踩坑了。咱们一步步理清楚:
你看,计算Accuracy的时候,你把self.input_y用tf.argmax(self.input_y, 1)转换成了整数类型的类别索引(比如0、1、2这类值),再和预测结果对比,这完全没问题。但到了计算Precision的时候,你直接传了self.input_y(这是one-hot编码的向量)和self.logits(模型的原始输出,不是类别索引),这俩参数的格式根本不符合tf.metrics.precision的要求。
划重点:tf.metrics.precision的参数规则
- 第一个参数
labels:必须是整数类别索引,不能是one-hot向量; - 第二个参数
predictions:得是和labels同格式的预测类别索引(多分类场景下),如果是二分类也可以传概率值,但多分类必须传索引。
修正后的代码
把参数转换成正确格式就行:
with tf.name_scope("pointwise_accuracy"): correct_predictions = tf.equal(self.predictions, tf.argmax(self.input_y, 1)) self.classification_accuracy = tf.reduce_mean(tf.cast(correct_predictions, "float"), name="accuracy") # 将one-hot标签转为类别索引,将logits转为预测类别索引 true_labels = tf.argmax(self.input_y, 1) pred_labels = tf.argmax(self.logits, 1) self.precision = tf.metrics.precision(true_labels, pred_labels, name="precision")[0]
额外提醒
别忘了在启动训练或评估前,初始化TensorFlow的局部变量(因为tf.metrics系列函数会创建状态变量来累计计算指标),比如在会话里执行:
sess.run(tf.local_variables_initializer())
这样调整后,Precision应该就能正常输出正确结果啦~
内容的提问来源于stack exchange,提问作者Abhay Singh
相关产品推荐
相关产品推荐

