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

Keras加载Transformer ASR训练模型时出现大量层变量加载失败的问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.07 10:24:34