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

跨环境加载自定义Keras机器翻译模型触发ValueError问题

自定义Keras机器翻译模型加载时的嵌入层变量错误问题

问题详情

我制作了一个用于机器翻译的自定义Keras模型,保存为tf_model.keras后,在非训练环境加载时触发ValueError,但在训练模型的Colab文件中保存并加载是成功的。

已做的处理:

  • 在自定义模型类顶部添加了@keras.saving.register_keras_serializable()装饰器
  • 实现了get_config()和from_config()方法

加载时的错误信息:

ValueError: Layer 'embedding' expected 1 variables, but received 0 variables during loading. Expected: ['embedding/embeddings:0']

模型代码:

@keras.saving.register_keras_serializable()
class BidirectionalEncoderandDecoderWithAttention(keras.Model):
    def __init__(**kwargs):
        super().__init__(**kwargs)
        self.encoder_embedding = Embedding(14617, 256, mask_zero=True)
        self.encoder = Bidirectional(LSTM(512 // 2, return_sequences=True, return_state=True))
        self.decoder_embedding = Embedding(29604, 256, mask_zero=True)
        self.decoder = LSTM(512, return_sequences=True)
        self.attention = Attention()
        self.output_layer = Dense(29604, activation='softmax')

    def call(self, inputs):
        encoder_inputs, decoder_inputs = inputs
        encoder_embeddings = self.encoder_embedding(encoder_inputs)
        decoder_embeddings = self.decoder_embedding(decoder_inputs)

        encoder_op, *encoder_state = self.encoder(encoder_embeddings)
        encoder_state = [
            tf.concat(encoder_state[0::2], axis=-1),
            tf.concat(encoder_state[1::2], axis=-1),
        ]
        decoder_op = self.decoder(decoder_embeddings, initial_state=encoder_state)
        attention_output = self.attention([decoder_op, encoder_op])
        output = self.output_layer(attention_output)
        return output

    def get_config(self):
        config = super().get_config()
        return config

    @classmethod
    def from_config(cls, config):
        return cls(**config)

加载代码:

my_model = keras.models.load_model("tf_model.keras", custom_objects={ 'BidirectionalEncoderandDecoderWithAttention': BidirectionalEncoderandDecoderWithAttention})

注:加载时已在同一文件中定义了BidirectionalEncoderandDecoderWithAttention类。

解决方案

问题出在自定义模型的序列化逻辑上,当前的get_config()方法没有保存子层的配置信息,导致加载时无法正确重建嵌入层等子层的变量。

修改步骤:

  1. 完善get_config()方法:
    在父类的config基础上,添加所有子层的配置,确保序列化时能保存子层的参数。修改后的get_config()如下:

    def get_config(self):
        config = super().get_config()
        # 保存所有子层的配置
        config.update({
            "encoder_embedding": keras.layers.serialize(self.encoder_embedding),
            "encoder": keras.layers.serialize(self.encoder),
            "decoder_embedding": keras.layers.serialize(self.decoder_embedding),
            "decoder": keras.layers.serialize(self.decoder),
            "attention": keras.layers.serialize(self.attention),
            "output_layer": keras.layers.serialize(self.output_layer),
        })
        return config
    
  2. 修改from_config()方法:
    从配置中反序列化所有子层,并初始化到模型实例中:

    @classmethod
    def from_config(cls, config):
        # 提取子层配置并反序列化
        encoder_embedding = keras.layers.deserialize(config.pop("encoder_embedding"))
        encoder = keras.layers.deserialize(config.pop("encoder"))
        decoder_embedding = keras.layers.deserialize(config.pop("decoder_embedding"))
        decoder = keras.layers.deserialize(config.pop("decoder"))
        attention = keras.layers.deserialize(config.pop("attention"))
        output_layer = keras.layers.deserialize(config.pop("output_layer"))
        
        # 创建模型实例
        model = cls(**config)
        # 赋值子层属性
        model.encoder_embedding = encoder_embedding
        model.encoder = encoder
        model.decoder_embedding = decoder_embedding
        model.decoder = decoder
        model.attention = attention
        model.output_layer = output_layer
        return model
    
  3. 重新保存模型:
    使用修改后的模型类重新训练并保存,或者在原训练环境中加载旧模型后用新的序列化逻辑重新保存:

    # 假设原模型已加载
    model.save("tf_model_fixed.keras")
    

额外说明

在Colab中能正常加载是因为训练环境中模型的子层变量还在内存中,序列化时隐式保留了相关信息,但跨环境加载时必须显式保存子层配置才能正确重建模型。

内容的提问来源于stack exchange,提问作者user8916969

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 04:40:15