讲师中心 微信公众号
AI工具推荐 视频效率加速

如何使用Python在TensorFlow中实现自定义的评价指标(Metrics)?

千墨酱_3665

千墨酱_3665

发布时间:2026-08-30 09:56:24

|

783人浏览过

|

来源于php中文网

原创

TensorFlow 2.x 自定义 metrics 必须继承 tf.keras.metrics.Metric 类并用 self.add_weight() 管理状态,普通函数无法累积批次结果;macro-F1 需按类别分别统计 TP/FP/FN 再均值,禁用 tf.py_function 生产环境使用。

如何使用python在tensorflow中实现自定义的评价指标(metrics)?

TensorFlow 2.x 中自定义 metrics 必须继承 tf.keras.metrics.Metric

直接写个普通函数(比如 def my_f1(y_true, y_pred))无法被 model.compile() 正确识别为指标——它不会累积批次结果,也无法在验证时自动重置状态。TensorFlow 的 metrics 是有状态的对象,必须通过类封装内部变量(如 self.tp, self.fn),并在 update_state() 和 result() 中明确定义更新与读取逻辑。

常见错误是试图用 tf.reduce_mean 或 tf.math.confusion_matrix 一次性计算整个 batch 的指标值后返回标量,这会导致训练/验证过程中指标无法跨 batch 累积(比如 precision 分母漏加其他 batch 的 false positives)。

实操建议:

  • 子类必须调用 super().__init__() 初始化父类状态管理机制
  • 所有可累积的中间变量(如 TP、FP、FN)必须用 self.add_weight() 创建,不能用 Python 普通变量或 tf.Variable 手动初始化
  • update_state() 接收的是未展平的 batch 张量(y_true 形状可能是 (batch_size, num_classes)),需自行做 argmax / threshold 判断
  • 若指标依赖预测概率(如 AUC),注意 y_pred 通常已是 softmax 输出,无需再套 sigmoid

实现多分类 F1-score(macro)的关键:分通道统计 + 最终取均值

macro-F1 要求对每个类别单独算 F1,再算平均,不能直接对全局 TP/FP/FN 计算。这意味着你得为每个类别维护一组 tp, fp, fn 变量——num_classes 决定了 add_weight 的 shape 参数。

立即学习“Python免费学习笔记(深入)”;

示例中容易踩的坑:

  • 忘记对 y_true 做 tf.one_hot 或 tf.argmax 对齐维度,导致布尔掩码错位
  • 用 tf.cast(y_pred > 0.5, tf.int32) 处理多分类输出(应改用 tf.argmax(y_pred, axis=-1))
  • 在 result() 里对分母为 0 的类别未做 tf.where 防御,导致 NaN 传播

简化版 macro-F1 核心片段:

Li Python Sec Check
Li Python Sec Check

Python 安全规范检查工具:基于 CloudBase 规范、腾讯安全指南,LLM 智能分析(默认禁用,优先本地执行)

下载
class MacroF1(tf.keras.metrics.Metric):
    def __init__(self, num_classes, name='macro_f1', **kwargs):
        super().__init__(name=name, **kwargs)
        self.num_classes = num_classes
        self.tp = self.add_weight(name='tp', shape=(num_classes,), initializer='zeros')
        self.fp = self.add_weight(name='fp', shape=(num_classes,), initializer='zeros')
        self.fn = self.add_weight(name='fn', shape=(num_classes,), initializer='zeros')
<pre class='brush:python;toolbar:false;'>def update_state(self, y_true, y_pred, sample_weight=None):
    y_true = tf.argmax(y_true, axis=-1)
    y_pred = tf.argmax(y_pred, axis=-1)
    for i in range(self.num_classes):
        tp_mask = tf.logical_and(tf.equal(y_true, i), tf.equal(y_pred, i))
        fp_mask = tf.logical_and(tf.not_equal(y_true, i), tf.equal(y_pred, i))
        fn_mask = tf.logical_and(tf.equal(y_true, i), tf.not_equal(y_pred, i))
        self.tp[i].assign_add(tf.reduce_sum(tf.cast(tp_mask, tf.float32)))
        self.fp[i].assign_add(tf.reduce_sum(tf.cast(fp_mask, tf.float32)))
        self.fn[i].assign_add(tf.reduce_sum(tf.cast(fn_mask, tf.float32)))

def result(self):
    f1_per_class = tf.zeros(self.num_classes)
    for i in range(self.num_classes):
        precision = self.tp[i] / (self.tp[i] + self.fp[i] + 1e-6)
        recall = self.tp[i] / (self.tp[i] + self.fn[i] + 1e-6)
        f1_per_class = tf.tensor_scatter_nd_update(
            f1_per_class,
            [[i]],
            [2 * precision * recall / (precision + recall + 1e-6)]
        )
    return tf.reduce_mean(f1_per_class)

使用 tf.py_function 包裹 scikit-learn 指标要格外小心

虽然可以用 tf.py_function 把 sklearn.metrics.f1_score 塞进去,但会破坏图执行(graph mode),导致无法保存 SavedModel、XLA 加速失效,且在 TPU 上直接报错。仅限 debug 阶段临时验证逻辑,绝不可用于生产训练循环。

更隐蔽的问题是数据类型和形状不匹配:tf.py_function 输入张量默认是 tf.int64,而 sklearn 多数函数要求 numpy array + int32 或 float64;若未显式指定 Tout 或在函数内做 .numpy().astype() 转换,会静默出错或返回错误值。

替代方案更稳妥:

  • 优先用原生 TensorFlow ops 重写逻辑(如上面的 MacroF1)
  • 若必须用 sklearn,只在 model.evaluate() 后对全量预测结果调用(即 CPU 上离线计算),而非作为 metrics= 参数传入
  • 避免在 update_state() 中调用 tf.py_function —— 它无法跨 batch 维持 sklearn 内部状态(如 classification_report 的计数器)

自定义 metric 在 model.compile() 和回调中的行为差异

同一个自定义 metric 类实例,在 compile(metrics=[MyMetric()]) 时会被框架自动复用并共享状态;但若在 tf.keras.callbacks.Callback 里手动创建新实例(如 on_test_batch_end 中 new 一个),其状态完全独立,数值无意义。

另一个易忽略点:reset_states() 不仅会在每个 epoch 开始前被调用,也会在 evaluate() 开始前触发。如果你在 update_state() 里做了副作用操作(如写文件、发 HTTP 请求),必须确保它们不会因重复 reset 而异常触发。

调试建议:

  • 在 update_state() 开头加 print("update_state called with", y_true.shape)(注意仅限 eager mode,否则 print 不生效)
  • 把自定义 metric 实例赋值给变量(如 my_f1 = MacroF1(3)),然后在训练后直接调用 my_f1.result().numpy() 查看当前值,比只看日志更可靠
  • 如果指标值始终为 0 或 NaN,先检查 add_weight 的 shape 是否与实际类别数一致,再确认 update_state 中是否真的执行了 assign_add

TensorFlow 自定义 metrics 的核心约束在于「状态必须由框架托管」,任何绕过 add_weight + assign_add 的累加方式,都会在分布式训练或多 epoch 场景下失效。别想着省事用 Python list append,那只会让你在验证集上看到完全不可信的数字。

热门AI工具

更多
DeepSeek

DeepSeek是一款面向对话、写作、编程和推理场景的AI大模型工具。

Laper
Laper Hot

Laper是专为编剧、导演和制片人推出的 AI 原生剧本创作工具。

音述AI
音述AI Hot

一款AI音频处理工具,主要用于音述AI是一个以“用声音述说故事”为核心的 AI 音乐创作与声音分享社区,适合需要提升相关任务效率的用户。

立刻MV
立刻MV Hot

立刻MV是一款AI文本写作工具,AI 音乐视频(MV)创作工具。

WorkBuddy

一款AI办公效率工具,主要用于腾讯云推出的AI原生桌面智能体工作台,适合需要提升相关任务效率的用户。

豆包大模型

豆包大模型是一款由字节跳动推出的企业级大语言模型服务平台。

AionClaw
AionClaw Hot

AionClaw是一款面向办公、创作和编程任务的AI桌面智能体。

蛙蛙写作

一款AI论文写作工具,主要用于超级AI智能写作助手,适合需要提升相关任务效率的用户。

SkildArt
SkildArt Hot

SkildArt是一款AI文本写作工具,一站式 AI 视觉创作平台。

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

1671

2023.07.20

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

4224

2023.07.25

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

1669

2023.07.31

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

24457

2023.08.03

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2987

2023.08.04

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

3027

2023.08.04

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

1163

2023.08.11

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

596

2023.08.10

FrankenPHP集成Laravel详细教程
FrankenPHP集成Laravel详细教程

本专题提供FrankenPHP集成Laravel的详细配置指南,全面解析运行原理、开发环境搭建、Caddyfile配置、Octane工作模式、数据库连接、队列任务、定时任务和生产环境优化,解决部署过程中常见的报错与兼容性问题。

40

2026.10.08

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程
关于我们 免责申明 举报中心 意见反馈 讲师合作 广告合作 最新更新
php中文网:公益在线php培训,帮助PHP学习者快速成长!
关注服务号
PHP中文网订阅号
每天精选资源文章推送

Copyright 2014-2026 https://www.php.cn/ All Rights Reserved | php.cn