加载Keras手写识别训练模型报Unknown layer: Custom>CTCLayer错误如何解决
问题解决步骤
- 第一步:修正CTCLayer类的实现,自定义Keras层必须实现
get_config方法才能正常序列化/反序列化,同时你当前的__init__方法对name参数的处理不符合Keras规范,修正后的代码如下:
import tensorflow as tf from tensorflow import keras class CTCLayer(keras.layers.Layer): def __init__(self, name=None, **kwargs): # 不要自行赋值self.name,交给父类初始化处理 super().__init__(name=name, **kwargs) self.loss_fn = keras.backend.ctc_batch_cost def call(self, y_true, y_pred): batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64") input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64") label_length = tf.cast(tf.shape(y_true)[1], dtype="int64") input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64") label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64") loss = self.loss_fn(y_true, y_pred, input_length, label_length) self.add_loss(loss) # 测试阶段直接返回预测结果 return y_pred # 新增get_config方法,用于层的序列化 def get_config(self): config = super().get_config() # 如果有自定义初始化参数要在这里添加,当前CTCLayer没有额外参数,直接返回父类config即可 return config
- 第二步:统一依赖导入,不要混用独立Keras包和TensorFlow内置的Keras,删掉冗余的
from keras.models import load_model导入语句,所有Keras相关接口都从tensorflow.keras下调用,避免版本冲突导致自定义层识别失败。 - 第三步:使用自定义对象作用域加载模型,比直接传
custom_objects参数兼容性更强,加载代码如下:
from tensorflow import keras with keras.utils.custom_object_scope({'CTCLayer': CTCLayer}): model = keras.models.load_model("test4_20211113.h5")
如果以上操作还是报错,可以给CTCLayer类加上Keras序列化注册装饰器,在类定义上方加@keras.utils.register_keras_serializable()即可,无需修改其他代码。
内容的提问来源于stack exchange,提问作者Hiroji
相关产品推荐
相关产品推荐

