调用tf.keras.models.clone_model出现自定义初始化器Unknown initializer错误
问题原因
tf.keras.models.clone_model的内部序列化逻辑和模型创建时存在差异:
- 构建模型时传入字符串
"complex_glorot_uniform",Keras会直接匹配全局自定义对象的对应键完成初始化 - 克隆模型时,Keras会先提取已初始化层的
kernel_initializer实例,序列化时默认用类名ComplexGlorotUniform作为检索键,你仅注册了字符串标识对应的实例,没有注册类名对应的实现,因此触发报错。
解决方法
方法1:补全全局自定义对象注册项
同时注册字符串标识和类名,无需修改克隆逻辑:
from tensorflow.keras.utils import get_custom_objects # 同时注册字符串键和类名键 init_dispatcher = { "complex_glorot_uniform": ComplexGlorotUniform(), "ComplexGlorotUniform": ComplexGlorotUniform } get_custom_objects().update(init_dispatcher)
方法2:克隆时显式传入自定义对象
不需要修改全局注册逻辑,调用克隆方法时单独传入配置即可:
cloned_model = tf.keras.models.clone_model( original_model, custom_objects={"ComplexGlorotUniform": ComplexGlorotUniform} )
补充建议
自定义初始化器最好实现get_config和from_config方法,避免后续序列化/反序列化场景出现异常:
class ComplexGlorotUniform(Initializer): # 你的原有初始化逻辑 def get_config(self): # 若初始化有参数可在此处返回对应键值对 return {} @classmethod def from_config(cls, config): return cls(**config)
内容的提问来源于stack exchange,提问作者J Agustin Barrachina
相关产品推荐
相关产品推荐

