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

tf.constant与tf.placeholder在tf.metrics封装中表现差异问题

问题分析与解决方案

你的问题出在对tf.metrics的使用逻辑误解,以及当前Sonnet模块实现的设计缺陷上。

为什么test3正常,test4返回0.0?

tf.metrics系列函数(比如tf.metrics.accuracy)的核心设计目的是累计计算整个数据集的性能指标,而非单批次即时指标。它会创建局部变量(比如total_correct和total_count)来累计所有批次的正确数和样本数:

  • 返回的第一个值metric是累计的平均值
  • 第二个值metric_update是更新这些累计变量的操作

在test3中,labels和logits是常量张量:

  1. 执行sess.run(accuracy)时,tf.control_dependencies会先触发metric_update,一次性计算整个常量数据集的正确数和样本数,直接更新累计变量
  2. 随后读取累计变量的值,自然得到整个数据集的准确率0.88,结果正确

而test4使用placeholder时,虽然逻辑上等价,但由于TensorFlow的执行优化或Sonnet模块的作用域管理问题,导致metric_update没有正确触发变量更新,最终读取到了局部变量初始化后的初始值0。更关键的是,这个实现本身不符合你“衡量每个批次性能”的需求——tf.metrics是累计指标,不是单批次指标。

正确实现:计算单批次性能指标

如果你需要的是每个批次的即时性能,应该直接通过TensorFlow基础操作计算,而非依赖tf.metrics。以下是修改后的Sonnet模块:

import tensorflow as tf
import sonnet as snt

class Metrics(snt.AbstractModule):
    def __init__(self, indicator, summaries=None, name="metrics"):
        super(Metrics, self).__init__(name=name)
        self._indicator = indicator
        self._summaries = summaries
        # 防止除零错误的极小值
        self._epsilon = tf.constant(1e-7, dtype=tf.float32)

    def _build(self, labels, logits):
        # 统一处理:如果logits是原始预测值(float类型),先转成类别索引
        predictions = logits if logits.dtype == labels.dtype else tf.argmax(logits, axis=-1)
        predictions = tf.cast(predictions, tf.int32)
        labels = tf.cast(labels, tf.int32)

        if self._indicator == "accuracy":
            correct_preds = tf.equal(labels, predictions)
            outputs = tf.reduce_mean(tf.cast(correct_preds, tf.float32))
        
        elif self._indicator == "precision":
            true_pos = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 1), tf.equal(predictions, 1)), tf.float32))
            false_pos = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 0), tf.equal(predictions, 1)), tf.float32))
            outputs = true_pos / (true_pos + false_pos + self._epsilon)
        
        elif self._indicator == "recall":
            true_pos = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 1), tf.equal(predictions, 1)), tf.float32))
            false_neg = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 1), tf.equal(predictions, 0)), tf.float32))
            outputs = true_pos / (true_pos + false_neg + self._epsilon)
        
        elif self._indicator == "f1_score":
            true_pos = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 1), tf.equal(predictions, 1)), tf.float32))
            false_pos = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 0), tf.equal(predictions, 1)), tf.float32))
            false_neg = tf.reduce_sum(tf.cast(tf.logical_and(tf.equal(labels, 1), tf.equal(predictions, 0)), tf.float32))
            
            precision = true_pos / (true_pos + false_pos + self._epsilon)
            recall = true_pos / (true_pos + false_neg + self._epsilon)
            
            outputs = 2 * precision * recall / (precision + recall + self._epsilon)
        
        else:
            raise ValueError(f"Unsupported metric: {self._indicator}")

        if isinstance(self._summaries, list):
            self._summaries.append(tf.summary.scalar(self._indicator, outputs))
        
        return outputs

修改后的优势

  1. 单批次即时计算:直接针对当前输入批次计算性能,完全符合你“衡量每个批次模型性能”的需求
  2. 无累计变量依赖:不需要管理局部变量,也无需初始化局部变量,彻底避免了原实现中的变量更新问题
  3. 鲁棒性更强:添加epsilon防止除零错误,同时兼容logits是原始预测值(float类型)的场景

验证修改后的代码

无论是test3还是test4的代码,使用这个修改后的Metrics模块都能正确返回结果。比如运行test4:

def test4():
    import numpy as np
    tf_labels = tf.placeholder(dtype=tf.int32, shape=[None])
    tf_logits = tf.placeholder(dtype=tf.int32, shape=[None])
    labels = np.array([1, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0], np.int32)
    logits = np.array([1, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0], np.int32)
    
    metrics = Metrics("accuracy")
    accuracy = metrics(tf_labels, tf_logits)
    metrics2 = Metrics("f1_score")
    f1_score = metrics2(tf_labels, tf_logits)
    
    writer = tf.summary.FileWriter("utils-feed", tf.get_default_graph())
    with tf.Session() as sess:
        sess.run(tf.global_variables_initializer())
        accu, f1 = sess.run([accuracy, f1_score], feed_dict={tf_labels: labels, tf_logits: logits})
        print(accu)  # 输出0.88
        print(f1)    # 输出约0.6667
        writer.close()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 04:23:16