Keras加载训练模型提示Unknown metric function: lr报错如何解决
错误原因
- 第一次报错是因为模型编译时使用了自定义的
lr指标,Keras保存模型时不会自带自定义对象的实现,加载时找不到对应定义触发报错 - 第二次报错是因为你在加载脚本中没有定义
lr这个函数变量,直接传未定义的变量自然触发NameError
解决方案
根据你加载模型的用途选择对应方案即可:
场景1:加载模型仅用于推理,不需要继续训练
直接在加载时添加compile=False参数,跳过模型编译步骤,不会检查自定义指标,代码如下:
from tensorflow import keras import os model_dir = 'My Directory' model1 = os.path.join(model_dir, "DenseNet_model_keras.h5") # 跳过编译,直接加载权重和结构用于推理 Vgg16 = keras.models.load_model(model1, compile=False)
场景2:加载模型后需要继续训练
需要在加载脚本中完整复现训练时的自定义对象定义,再传入custom_objects参数:
from tensorflow import keras from tensorflow.keras.optimizers import RMSprop import os # 1. 复现自定义lr指标的定义 opt = RMSprop() def get_lr_metric(optimizer): def lr(y_true, y_pred): return optimizer.lr return lr lr_track = get_lr_metric(opt) # 2. 如果你后续训练还要用到CyclicLR回调,需要把你从GitHub克隆的CyclicLR类完整实现复制到此处 # 3. 加载模型时传入自定义对象 model_dir = 'My Directory' model1 = os.path.join(model_dir, "DenseNet_model_keras.h5") Vgg16 = keras.models.load_model(model1, custom_objects={ "lr": lr_track })
内容的提问来源于stack exchange,提问作者maryam gh
相关产品推荐
相关产品推荐

