You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

胶囊网络模型保存加载时自定义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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.25 07:06:04