训练后无法加载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
相关产品推荐
相关产品推荐

