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

TensorFlow U-Net自定义average_accuracy在fit与evaluate结果差异过大问题

根因分析
  • 核心问题是你使用普通Python函数作为Keras自定义指标,不符合Keras指标的设计规则。Keras在fit阶段对训练集指标的计算逻辑是:每个batch独立调用指标函数得到该batch的指标值,所有batch的结果取均值作为epoch的指标输出;而evaluate阶段是对整个数据集做全局计算,两者逻辑完全不一致。你的逐类准确率取均值的计算逻辑对数据分布非常敏感,batch级平均和全局平均的结果会存在巨大差异。
  • 自定义指标中的动态逻辑(动态TensorArray、if判断样本是否存在)在无状态函数模式下,无法跨batch累计每类的正确数和总数,进一步放大了batch级计算和全局计算的误差。
  • 数据集shuffle导致的结果差异:你的train_ds默认带shuffle操作,evaluate调用时会迭代shuffle后的样本,每个batch的类别组成随机,而转成numpy数组传入时是固定的全局样本,两者的输入分布不一致导致结果不同;取消shuffle后样本顺序和batch组成固定,差异自然消失。
  • 验证集结果一致是巧合:验证集的类别分布均匀、batch大小合适,刚好batch级平均的结果和全局计算结果接近,不代表指标实现正确。
解决方案

1. 重写自定义指标为有状态类

继承tf.keras.metrics.Metric实现指标,跨batch累计每类的总样本数和正确样本数,最终统一计算全局的平均准确率,示例实现如下:

import tensorflow as tf

class AverageAccuracy(tf.keras.metrics.Metric):
    def __init__(self, num_classes=4, name='average_accuracy', **kwargs):
        super().__init__(name=name, **kwargs)
        self.num_classes = num_classes
        # 初始化累计变量:每类正确数、每类总样本数
        self.cls_correct = self.add_weight(
            name='cls_correct', shape=(num_classes,), initializer='zeros'
        )
        self.cls_total = self.add_weight(
            name='cls_total', shape=(num_classes,), initializer='zeros'
        )

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 过滤全0标签
        remove_zeros_mask = tf.math.logical_not(
            tf.math.reduce_all(tf.math.logical_not(tf.cast(y_true, bool)), axis=-1)
        )
        y_true = tf.boolean_mask(y_true, remove_zeros_mask)
        y_pred = tf.boolean_mask(y_pred, remove_zeros_mask)
        # 转类别ID
        y_true = tf.argmax(y_true, axis=-1)
        y_pred = tf.argmax(y_pred, axis=-1)
        # 逐类更新累计值
        for cls_id in range(self.num_classes):
            cls_mask = y_true == cls_id
            cls_total = tf.reduce_sum(tf.cast(cls_mask, tf.float32))
            if cls_total > 0:
                cls_correct = tf.reduce_sum(
                    tf.cast(tf.logical_and(cls_mask, y_pred == cls_id), tf.float32)
                )
                self.cls_correct[cls_id].assign_add(cls_correct)
                self.cls_total[cls_id].assign_add(cls_total)

    def result(self):
        # 仅统计出现过的类的准确率,再取均值
        exist_cls_mask = self.cls_total > 0
        exist_cls_acc = tf.boolean_mask(self.cls_correct / self.cls_total, exist_cls_mask)
        return tf.reduce_mean(exist_cls_acc)

    def reset_state(self):
        # 每个epoch/evaluate前重置累计变量
        self.cls_correct.assign(tf.zeros_like(self.cls_correct))
        self.cls_total.assign(tf.zeros_like(self.cls_total))

使用时直接在模型编译阶段传入该类的实例即可:

model.compile(
    optimizer='adam',
    loss='categorical_crossentropy',
    metrics=['accuracy', AverageAccuracy(num_classes=4)]
)

2. 单独定义评估用训练集

确保evaluate使用的训练集不带shuffle、repeat操作,避免采样样本和原始训练集不一致:

# 移除shuffle、repeat操作,仅保留分批和预处理逻辑
eval_train_ds = train_ds.unbatch().batch(batch_size) # 或者按原数据流水线重新定义无shuffle版本

3. 验证BN层状态

Keras默认在evaluate阶段会自动将模型切换为推理模式,BN层使用滑动均值/方差计算,无需额外调整,可排除BN层导致的巨大差异。

内容的提问来源于stack exchange,提问作者J Agustin Barrachina

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 13:27:04