基于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
相关产品推荐
相关产品推荐

