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

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 参数。
示例中容易踩的坑:
- 忘记对
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 核心片段:
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,那只会让你在验证集上看到完全不可信的数字。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











