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

基于TensorFlow图像分割教程添加自定义指标保存模型报错咨询

报错原因

你使用了自定义的precision评估指标,TensorFlow序列化保存模型时不会自动存储自定义函数的实现逻辑,加载模型时找不到对应指标定义就会触发该报错。

解决方案

方法一:加载时指定自定义对象(快速修复,适合简单函数式指标)

你现有的模型保存代码无需修改,加载模型时传入custom_objects参数,将自定义指标的名称和对应的函数绑定即可:

from tensorflow.keras.models import load_model
from tensorflow.keras import backend as K

# 注意:加载脚本里必须先定义好precision函数,再执行加载
def precision(y_true, y_pred):
    y_true = K.ones_like(y_true) 
    true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))
    predicted_positives = K.sum(K.round(K.clip(y_pred, 0, 1)))
    precision = true_positives / (predicted_positives + K.epsilon())
    return precision

model = load_model('mymodel', custom_objects={'precision': precision})

参数里的键precision必须和你编译模型时传入metrics的函数名完全一致。

方法二:封装为标准Metric类(稳定方案,适合复用/复杂指标)

将自定义指标继承tf.keras.metrics.Metric基类,实现序列化所需的get_config和from_config方法,兼容官方序列化规则:

import tensorflow as tf
from tensorflow.keras import backend as K

# 可选:添加序列化注册装饰器,后续加载无需传custom_objects
@tf.keras.utils.register_keras_serializable()
class Precision(tf.keras.metrics.Metric):
    def __init__(self, name='precision', **kwargs):
        super().__init__(name=name, **kwargs)
        self.true_positives = self.add_weight(name='tp', initializer='zeros')
        self.predicted_positives = self.add_weight(name='pp', initializer='zeros')

    def update_state(self, y_true, y_pred, sample_weight=None):
        y_true = K.ones_like(y_true)
        true_pos = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))
        pred_pos = K.sum(K.round(K.clip(y_pred, 0, 1)))
        self.true_positives.assign_add(true_pos)
        self.predicted_positives.assign_add(pred_pos)

    def result(self):
        return self.true_positives / (self.predicted_positives + K.epsilon())

    def reset_state(self):
        self.true_positives.assign(0.)
        self.predicted_positives.assign(0.)

    def get_config(self):
        base_config = super().get_config()
        return base_config
    
    @classmethod
    def from_config(cls, config):
        return cls(**config)

编译模型时传入类的实例即可:

model.compile(optimizer='adam',
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy', Precision()])

添加了register_keras_serializable装饰器后,加载模型无需额外传参,直接调用load_model('mymodel')即可。

注意

你现有precision函数里强制将y_true设为全1的逻辑需要自行验证是否符合你的业务场景,以上方案仅解决模型保存加载的报错问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 07:39:03