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

训练后无法加载Vision Transformer模型的问题求助

解决Keras Vision Transformer模型保存加载的TypeError问题

问题根源

官网示例里的Patches是自定义Keras层,它的__init__方法没显式接收name参数,但Keras保存模型时会自动把层的名称写入配置,加载时会尝试把这个name参数传给__init__,就触发了参数不匹配的报错。

具体解决步骤

1. 修复自定义层的定义

把Patches类的__init__改成支持name参数,同时确保get_config能正确导出配置:

class Patches(layers.Layer):
    def __init__(self, patch_size, name=None):
        super().__init__(name=name)  # 把name传给父类Layer
        self.patch_size = patch_size

    def call(self, images):
        # 保留官网示例的call方法逻辑
        batch_size = tf.shape(images)[0]
        patches = tf.image.extract_patches(
            images=images,
            sizes=[1, self.patch_size, self.patch_size, 1],
            strides=[1, self.patch_size, self.patch_size, 1],
            rates=[1, 1, 1, 1],
            padding="VALID",
        )
        patch_dims = patches.shape[-1]
        patches = tf.reshape(patches, [batch_size, -1, patch_dims])
        return patches

    def get_config(self):
        config = super().get_config()
        config.update({"patch_size": self.patch_size})
        return config

另外,模型里的PatchEncoder自定义层也要做同样修改,避免后续加载时出同类问题:

class PatchEncoder(layers.Layer):
    def __init__(self, num_patches, projection_dim, name=None):
        super().__init__(name=name)
        self.num_patches = num_patches
        self.projection = layers.Dense(units=projection_dim)
        self.position_embedding = layers.Embedding(
            input_dim=num_patches, output_dim=projection_dim
        )

    def call(self, patch):
        positions = tf.range(start=0, limit=self.num_patches, delta=1)
        encoded = self.projection(patch) + self.position_embedding(positions)
        return encoded

    def get_config(self):
        config = super().get_config()
        config.update({"num_patches": self.num_patches, "projection_dim": self.projection_dim})
        return config

2. 重新训练并保存模型

用修改后的自定义层重新训练模型,再执行model.save("VITexp.keras")完成保存。

3. 在新Notebook中加载模型

在目标Notebook里,必须先定义好上述修改后的两个自定义层,之后再加载模型:

# 先复制Patches和PatchEncoder的完整定义代码到当前Notebook
from tensorflow.keras.models import load_model

# 方式1:先定义自定义层再加载
Lmodel = load_model("VITexp.keras")

# 方式2:如果不想提前定义,也可以用custom_objects参数指定
# Lmodel = load_model("VITexp.keras", custom_objects={"Patches": Patches, "PatchEncoder": PatchEncoder})

Lmodel.get_config()

关键提醒

  • 所有自定义层都要保证__init__接受name参数并传递给父类,同时get_config要包含所有自定义的参数
  • 如果模型用到了自定义损失函数或度量,加载时也要用同样的方式处理,要么提前定义,要么在load_model里通过custom_objects指定

内容的提问来源于stack exchange,提问作者B.RATH

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 18:07:38