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

带自定义TransformerEncoder层的Keras模型加载后权重异常问题

问题描述

我基于Francois Chollet的模板在Keras中实现了TransformerEncoder自定义层,训练完成后用model.save保存模型,但重新加载后推理时发现权重变回随机值,模型完全失去推理能力。我已经尝试了以下四种方案但均无效:

  • 在类上使用@tf.keras.utils.register_keras_serializable()装饰器
  • 在__init__中加入**kwargs
  • 为自定义层实现get_config和from_config方法
  • 使用custom_object_scope加载模型

可复现问题的最小代码:

import numpy as np
from tensorflow import keras
import tensorflow as tf
from tensorflow.keras import layers
from keras.models import load_model
from keras.utils import custom_object_scope

@tf.keras.utils.register_keras_serializable()
class TransformerEncoder(layers.Layer):
    def __init__(self, embed_dim, dense_dim, num_heads, **kwargs):
        super().__init__(**kwargs)
        self.embed_dim = embed_dim
        self.dense_dim = dense_dim
        self.num_heads = num_heads
        self.attention = layers.MultiHeadAttention(
            num_heads=num_heads, key_dim=embed_dim)
        self.dense_proj = keras.Sequential(
            [
                layers.Dense(dense_dim, activation="relu"),
                layers.Dense(embed_dim),
            ]
        )
        self.layernorm_1 = layers.LayerNormalization()
        self.layernorm_2 = layers.LayerNormalization()

    def call(self, inputs, mask=None):
        if mask is not None:
            mask = mask[:, tf.newaxis, :]
        attention_output = self.attention(
            inputs, inputs, attention_mask=mask)
        proj_input = self.layernorm_1(inputs + attention_output)
        proj_output = self.dense_proj(proj_input)
        return self.layernorm_2(proj_input + proj_output)

    def get_config(self):
        config = super().get_config()
        config.update({
            "embed_dim": self.embed_dim,
            "num_heads": self.num_heads,
            "dense_dim": self.dense_dim,
        })
        return config

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


# Create simple model:
encoder = TransformerEncoder(embed_dim=2, dense_dim=2, num_heads=1)
inputs = keras.Input(shape=(2, 2), batch_size=None, name="test_inputs")
x = encoder(inputs)
x = layers.Flatten()(x)
outputs = layers.Dense(1, activation="linear")(x)
model = keras.Model(inputs, outputs)

# Fit the model and save it:
np.random.seed(42)
X = np.random.rand(10, 2, 2)
y = np.ones(10)
model.compile(optimizer=keras.optimizers.Adam(), loss="mean_squared_error")
model.fit(X, y, epochs=2, batch_size=1)
model.save("./test_model")

# Load the saved model:
with custom_object_scope({
    'TransformerEncoder': TransformerEncoder
}):
    loaded_model = load_model("./test_model")

print(model.weights[0].numpy())
print(loaded_model.weights[0].numpy())
解决方案

问题核心是自定义层的子层在__init__中直接实例化时,Keras无法正确追踪子层的权重状态,导致保存和加载时子层权重被重置为初始随机值。

正确的修改方式是将子层的创建逻辑移到build方法中,该方法会在层接收输入形状后执行,Keras能在此过程中正确注册子层的权重并纳入保存/加载流程。修改后的TransformerEncoder类代码如下:

@tf.keras.utils.register_keras_serializable()
class TransformerEncoder(layers.Layer):
    def __init__(self, embed_dim, dense_dim, num_heads, **kwargs):
        super().__init__(**kwargs)
        self.embed_dim = embed_dim
        self.dense_dim = dense_dim
        self.num_heads = num_heads
        # 先初始化子层为None,在build中创建
        self.attention = None
        self.dense_proj = None
        self.layernorm_1 = None
        self.layernorm_2 = None

    def build(self, input_shape):
        # 在build方法中创建所有子层
        self.attention = layers.MultiHeadAttention(
            num_heads=self.num_heads, key_dim=self.embed_dim)
        self.dense_proj = keras.Sequential(
            [
                layers.Dense(self.dense_dim, activation="relu"),
                layers.Dense(self.embed_dim),
            ]
        )
        self.layernorm_1 = layers.LayerNormalization()
        self.layernorm_2 = layers.LayerNormalization()
        super().build(input_shape)

    def call(self, inputs, mask=None):
        if mask is not None:
            mask = mask[:, tf.newaxis, :]
        attention_output = self.attention(
            inputs, inputs, attention_mask=mask)
        proj_input = self.layernorm_1(inputs + attention_output)
        proj_output = self.dense_proj(proj_input)
        return self.layernorm_2(proj_input + proj_output)

    def get_config(self):
        config = super().get_config()
        config.update({
            "embed_dim": self.embed_dim,
            "num_heads": self.num_heads,
            "dense_dim": self.dense_dim,
        })
        return config

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

替换原代码中的TransformerEncoder类后,重新运行代码,保存和加载后的模型权重会完全一致,推理能力正常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 11:40:39