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

如何让Keras自定义指标可被ModelCheckpoint等回调正常识别使用

问题背景

自定义继承自ks.metrics.MeanIoU的SparseMeanIoU指标,可在model.fit训练阶段正常计算数值,也能被TensorBoard识别展示,但配置ModelCheckpoint监控该指标保存最优权重时,无论传入指标原名还是拼接val_前缀的名称,都会触发如下警告,无法执行保存逻辑:

WARNING:tensorflow:Can save best model only with Sparse_MeanIoU available, skipping.

相关代码如下:

自定义指标实现

N_CLASSES = 17
class SparseMeanIoU(ks.metrics.MeanIoU):
    def __init__(self,
                 y_true = None,
                 y_pred = None,
                 num_classes = None,
                 name = "Sparse_MeanIoU",
                 dtype = None):
        super(SparseMeanIoU, self).__init__(num_classes = num_classes,
                                             name = name, dtype = dtype)
        self.__name__ = "Sparse_MeanIoU"
    def get_config(self):
        return {"num_classes": self.num_classes, \
                "name": self.name, \
                    "dtype": self._dtype}
    def update_state(self, y_true, y_pred, sample_weight = None):
        y_pred = tf.math.argmax(y_pred, axis = -1)
        return super().update_state(y_true, y_pred, sample_weight)
    def __getstate__(self):
        variables = {v.name: v.numpy() for v in self.variables}
        state = { \
            name: variables[var.name] \
                for name, var in self._unconditional_dependency_names.items() \
                    if isinstance(var, tf.Variable)}
        state["name"] = self.name
        state["num_classes"] = self.num_classes
        return state
    def __setstate__(self, state):
        self.__init__(name = state.pop("name"), \
            num_classes = state.pop("num_classes"))
        for name, value in state.items():
            self._unconditional_dependency_names[name].assign(value)

met = SparseMeanIoU(num_classes = N_CLASSES)

回调配置

monitor_metric = met.name #"val_" + met.name
cptdir = "/some/directory/"
logdir = "/some/other/directory/"
cllbs = [
    ks.callbacks.ModelCheckpoint(os.path.join(cptdir, \
                                              "Epoch.{epoch:02d}.hdf5"), \
                                 monitor = monitor_metric, \
                                 mode = "max", \
                                 save_best_only = True, \
                                 save_freq = 5),
    ks.callbacks.TensorBoard(log_dir = logdir, histogram_freq = 5)
    ]

模型编译与训练配置

# 编译
model.compile(optimizer = "Adam", loss = ks.losses.SparseCategoricalCrossentropy(),
                 metrics = [met, "sparse_categorical_accuracy"])

# 训练
N_img = 100000 # 训练集图像数量
bs = 5 # 批次大小
args_fit = {"epochs" : 150,
            "steps_per_epoch" : np.ceil(N_img/bs),
            "validation_steps" : np.ceil(N_val/bs),
            "callbacks" : cllbs}
hist = model.fit(dataset["train"],
            validation_data = dataset["val"],
            **args_fit)
故障原因

问题由两个独立的配置错误共同导致:

  1. 保存触发逻辑与指标更新逻辑不匹配:ModelCheckpoint的save_freq参数设为整数5时,会每5个批次触发一次保存检查。但MeanIoU属于全epoch累计型指标,需要累计整个epoch所有样本的预测结果计算混淆矩阵后才能输出有效值,批次触发节点下该指标不会写入日志字典,回调自然找不到对应监控项。
  2. 自定义序列化逻辑破坏了Keras内部指标映射:手动重写的__getstate__/__setstate__方法没有调用父类实现,Keras编译模型时会对传入的指标做序列化拷贝,自定义序列化逻辑会导致拷贝生成的指标实例和实际计算写入日志的指标实例名称映射断裂,即便切换到epoch级触发也可能识别失败。

另外代码中冗余的self.__name__赋值、未调用父类的get_config实现也会提升名称匹配失败的概率。提到的on_epoch_end是TF1.x版本自定义指标的遗留接口,TF2.x下继承tf.keras.metrics.Metric实现的指标不需要额外实现该方法。

修复方案

按以下步骤调整代码即可让回调正常识别指标:

  1. 删除自定义的__getstate__、__setstate__方法,移除__init__中冗余的self.__name__赋值,get_config方法调用父类实现保证序列化完整性。
  2. 将ModelCheckpoint的save_freq改为"epoch",匹配累计型指标的更新周期;如果需要固定间隔epoch保存最优模型,可搭配自定义回调实现间隔判断,不要使用整数型批次级save_freq监控epoch级指标。
  3. 明确指定监控项为验证集指标val_Sparse_MeanIoU,不要动态取实例属性拼接名称,避免作用域导致的名称不匹配。

修复后的自定义指标代码:

N_CLASSES = 17
class SparseMeanIoU(ks.metrics.MeanIoU):
    def __init__(self,
                 num_classes,
                 name = "Sparse_MeanIoU",
                 dtype = None):
        super().__init__(num_classes=num_classes, name=name, dtype=dtype)
    
    def get_config(self):
        return super().get_config()

    def update_state(self, y_true, y_pred, sample_weight = None):
        y_pred = tf.math.argmax(y_pred, axis=-1)
        return super().update_state(y_true, y_pred, sample_weight)

met = SparseMeanIoU(num_classes=N_CLASSES)

修复后的回调配置:

cptdir = "/some/directory/"
logdir = "/some/other/directory/"
cllbs = [
    ks.callbacks.ModelCheckpoint(
        os.path.join(cptdir, "Epoch.{epoch:02d}.hdf5"),
        monitor = "val_Sparse_MeanIoU",
        mode = "max",
        save_best_only = True,
        save_freq = "epoch"
    ),
    ks.callbacks.TensorBoard(log_dir=logdir, histogram_freq=5)
]

校验方式:第一个epoch训练结束后,查看控制台打印的日志条目,确认存在val_Sparse_MeanIoU字段,ModelCheckpoint即可正常识别该指标并执行最优权重保存逻辑。

内容的提问来源于stack exchange,提问作者Manuel Popp

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 15:39:26