在Keras中使用tf.keras.metrics.AUC作为训练指标后,如何正确使用load_model加载模型?
解决加载模型时AUC指标识别失败的问题
这个问题我之前碰到过好几个类似的案例,核心坑点在于你在custom_objects里传入的AUC实例和训练时定义的参数不完全匹配,导致Keras无法对应上保存的metric配置。
问题根源
你训练时的AUC是带特定参数的:
model.compile(..., metrics=["accuracy", AUC(name="auc", curve="PR")])
这里明确指定了curve="PR"(计算PR曲线的AUC),但加载模型时你只传了AUC(name="auc"),少了curve="PR"这个关键参数。Keras保存模型时会完整记录metric的所有配置信息,加载时需要完全匹配的定义,否则就会抛出“Unknown metric function”的错误。
两种正确解决方法
方法一:传入AUC类(推荐,更简洁)
直接传入AUC的类,而不是实例,让Keras自动根据模型保存的参数(包括name="auc"和curve="PR")来重建metric:
from tensorflow.keras.metrics import AUC load_model(checkpoint, custom_objects={"auc": AUC})
这种方式最稳妥,不用手动复制所有参数,Keras会自己处理匹配逻辑。
方法二:传入参数完全一致的AUC实例
如果你一定要传实例,必须保证和训练时的参数完全相同,包括curve="PR":
from tensorflow.keras.metrics import AUC load_model(checkpoint, custom_objects={"auc": AUC(name="auc", curve="PR")})
这样实例的配置和训练时完全对齐,Keras就能正确识别对应的metric了。
额外提示
如果你的模型里还有其他自定义层或metric,都要遵循这个原则:要么传对应的类,要么传参数完全一致的实例,确保保存的配置和加载时的定义一一对应。
内容的提问来源于stack exchange,提问作者fmorenovr
相关产品推荐
相关产品推荐

