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

自定义TensorFlow模型保存后加载失败:TokenAndPositionEmbedding层序列化问题

解决TensorFlow自定义Embedding层加载失败问题

问题核心

用model.save保存含自定义TokenAndPositionEmbedding层的模型后,调用tf.keras.models.load_model时提示找不到该类,即使已经使用@keras.utils.register_keras_serializable装饰器。

解决方案

1. 统一装饰器与导入路径

不同TensorFlow/Keras版本中,序列化装饰器的路径有差异:

  • TensorFlow 2.9及更早版本:使用keras.utils.register_keras_serializable
  • TensorFlow 2.10+:推荐使用keras.saving.register_keras_serializable

关键注意点:加载模型前必须先执行自定义层的代码,确保Keras完成类的注册。

2. 加载模型时显式指定自定义对象

如果自动注册失效,直接在load_model中通过custom_objects参数传入自定义层类,强制Keras识别:

from tensorflow import keras
# 先定义或导入自定义层类
loaded_model = keras.models.load_model("你的模型保存路径", custom_objects={"TokenAndPositionEmbedding": TokenAndPositionEmbedding})

3. 修正自定义层代码细节

确保层内的layers导入来自正确的Keras模块,避免命名空间混乱:

from tensorflow import keras
from tensorflow.keras import layers
import tensorflow as tf

# 适配新版本的装饰器(TF2.10+)
@keras.saving.register_keras_serializable(package='Custom', name='TokenAndPositionEmbedding')
class TokenAndPositionEmbedding(keras.layers.Layer):
    def __init__(self, max_len, vocab_size, embed_dim, **kwargs):
        super().__init__(**kwargs)
        self.max_len = max_len
        self.vocab_size = vocab_size
        self.embed_dim = embed_dim
        self.token_emb = layers.Embedding(input_dim=vocab_size, output_dim=embed_dim)
        self.pos_emb = layers.Embedding(input_dim=max_len, output_dim=embed_dim)

    def call(self, x):
        maxlen = tf.shape(x)[-1]
        positions = tf.range(start=0, limit=maxlen, delta=1)
        positions = self.pos_emb(positions)
        x = self.token_emb(x)
        return x + positions

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

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

原因说明

  • 装饰器的注册逻辑需要在加载模型前触发,如果加载时自定义层代码未被执行,Keras无法读取到注册信息。
  • 不同版本的Keras对序列化装饰器的路径做了调整,混用路径会导致注册失效。
  • 显式指定custom_objects是最稳妥的方案,可绕过自动注册的潜在兼容性问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 15:22:39