胶囊网络模型保存加载时自定义squash函数未定义报错如何解决
胶囊网络自定义函数加载报错修复及保存方案
报错修复方法
squash函数找不到的核心原因是模型序列化保存时,仅将自定义函数的名称作为字符串引用存储,没有保存函数本身,可按以下步骤修复:
- 确保调用
load_model或model_from_json前,加载脚本中已经完整定义了所有用到的自定义函数(包括squash、margin_loss)和自定义层 - 给所有自定义函数加上Keras序列化注册装饰器,Keras会自动识别自定义对象,无需手动传入
custom_objects:
import tensorflow as tf from tensorflow.keras import backend as K @tf.keras.utils.register_keras_serializable() def squash(vectors, axis=-1): s_squared_norm = tf.reduce_sum(tf.square(vectors), axis, keep_dims=True) scale = s_squared_norm / (1 + s_squared_norm) / tf.sqrt(s_squared_norm + K.epsilon()) return tf.multiply(scale, vectors) # 同理给margin_loss也添加该装饰器 @tf.keras.utils.register_keras_serializable() def margin_loss(y_true, y_pred): # 你的margin_loss实现代码 pass
- 若你使用的是历史训练好的h5模型,无法修改训练时的代码,加载时需将所有自定义对象全部传入
custom_objects,不能遗漏:
model1 = load_model('model2.h5', custom_objects={ 'ClassCapsule':ClassCapsule, 'Length':Length, 'Mask':Mask, 'tf': tf, 'squash': squash, 'margin_loss': margin_loss # 必须补全所有自定义损失、函数 })
更稳定的复杂模型保存加载方案
针对带大量自定义层、函数的非序列模型,推荐使用TensorFlow官方的SavedModel格式替代h5/json格式,该格式会打包完整计算图、权重和自定义对象,加载时无需手动传入custom_objects,适配部署场景:
保存代码
# 训练完成后直接保存为SavedModel格式 model.save("capsule_model") # 生成同名文件夹,无需后缀
加载代码
# 无需传入任何自定义对象,直接加载即可使用 model = tf.keras.models.load_model("capsule_model")
内容的提问来源于stack exchange,提问作者Skander Zitoun
相关产品推荐
相关产品推荐

