Keras加载Transformer ASR训练模型时出现大量层变量加载失败的问题
看起来你遇到的核心问题是自定义Transformer相关组件的序列化/反序列化不完整,导致加载模型时无法正确重建子层(比如Dense、Embedding这些标准层)的变量。以下是几个针对性的解决方案,你可以逐一尝试:
1. 确保所有自定义类都实现了完整的get_config()方法
你已经给SpeechFeatureEmbedding写了get_config(),但其他自定义类(Transformer、TransformerEncoder、TransformerDecoder、TokenEmbedding、CustomSchedule)也需要同样处理。如果这些类的get_config()没有返回所有初始化参数,Keras在加载时就无法正确重建它们,进而导致其内部的标准层(比如报错的dense_65)参数不匹配。
举个Transformer类的示例:
@keras.saving.register_keras_serializable(package="ASRTransformer") class Transformer(keras.Model): def __init__(self, num_hid=200, num_head=2, num_feed_forward=400, target_maxlen=100, num_layers_enc=4, num_layers_dec=1, num_classes=34, **kwargs): super().__init__(**kwargs) # 保存初始化参数 self.num_hid = num_hid self.num_head = num_head self.num_feed_forward = num_feed_forward self.target_maxlen = target_maxlen self.num_layers_enc = num_layers_enc self.num_layers_dec = num_layers_dec self.num_classes = num_classes # ... 其他层初始化逻辑 def get_config(self): # 合并父类配置与自定义参数 base_config = super().get_config() custom_config = { "num_hid": self.num_hid, "num_head": self.num_head, "num_feed_forward": self.num_feed_forward, "target_maxlen": self.target_maxlen, "num_layers_enc": self.num_layers_enc, "num_layers_dec": self.num_layers_dec, "num_classes": self.num_classes, } return {**base_config, **custom_config} # ... 你的call方法实现
同样的,给TransformerEncoder、TransformerDecoder、TokenEmbedding都补上完整的get_config(),确保所有用于创建层的参数都被返回。
对于CustomSchedule(学习率调度器),也需要实现get_config():
@keras.saving.register_keras_serializable(package="ASRTransformer") class CustomSchedule(keras.optimizers.schedules.LearningRateSchedule): def __init__(self, d_model, warmup_steps=4000): super().__init__() self.d_model = d_model self.warmup_steps = warmup_steps def get_config(self): return { "d_model": self.d_model, "warmup_steps": self.warmup_steps, } def __call__(self, step): # ... 你的学习率计算逻辑
2. 使用custom_object_scope包裹保存和加载操作
虽然你已经在load_model中传入了custom_objects,但有时候用custom_object_scope上下文管理器包裹整个保存和加载流程会更可靠,确保Keras能正确识别所有自定义组件:
from keras.saving import custom_object_scope # 保存模型时 with custom_object_scope({ 'TokenEmbedding': TokenEmbedding, 'SpeechFeatureEmbedding': SpeechFeatureEmbedding, 'TransformerEncoder': TransformerEncoder, 'TransformerDecoder': TransformerDecoder, 'Transformer': Transformer, 'CustomSchedule': CustomSchedule }): model.save("model_1_epoch.keras") # 加载模型时 with custom_object_scope({ 'TokenEmbedding': TokenEmbedding, 'SpeechFeatureEmbedding': SpeechFeatureEmbedding, 'TransformerEncoder': TransformerEncoder, 'TransformerDecoder': TransformerDecoder, 'Transformer': Transformer, 'CustomSchedule': CustomSchedule }): model = keras.models.load_model("model_1_epoch.keras", compile=False)
3. 验证模型保存前已完全构建
虽然你已经训练了1个epoch,理论上模型已经被构建,但可以在保存前手动调用一次模型(用一个示例输入)确保所有层都已初始化:
# 替换成你实际的输入形状 dummy_input = tf.random.uniform((1, 100, 80)) dummy_target = tf.random.uniform((1, max_target_len), dtype=tf.int32) # 调用模型触发层构建 _ = model([dummy_input, dummy_target]) # 再保存模型 model.save("model_1_epoch.keras")
4. 坚持使用.keras格式
Keras的原生.keras格式比旧的.h5格式更稳定,尤其是对于包含自定义组件的模型。你之前尝试.h5也报错,所以继续使用.keras格式即可。
按照上面的步骤修复后,应该就能正确加载模型了。核心就是确保所有自定义组件的序列化逻辑完整,让Keras能精确重建训练时的模型结构。
内容来源于stack exchange

