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

TensorFlow 2.6加载含自定义层模型报pool1参数不识别错误

报错原因

你触发报错的核心原因是Localization类的get_config方法中,错误地将pool1、conv1等子层实例加入了序列化配置。
TensorFlow 2.6+版本对自定义层的参数校验逻辑收紧:加载模型时,get_config返回的所有键值对都会作为参数传入自定义层的__init__方法,父类tf.keras.layers.Layer无法识别pool1这类非内置参数,因此抛出参数不识别的TypeError。而2.5版本的校验规则更宽松,所以之前的流程可以正常运行。

修复方法

仅保留初始化需要的超参数到get_config的返回值即可,子层实例不需要写入配置:这些子层会在层初始化时根据超参数自动创建,模型权重也会在加载时自动恢复,无需额外序列化。

修正后的完整Localization类代码

class Localization(tf.keras.layers.Layer):
    def __init__(self, filters_1, filters_2, fc_units, kernel_size=(5,5), \
                 pool_size=(2,2), **kwargs):
        self.filters_1 = filters_1
        self.filters_2 = filters_2
        self.fc_units = fc_units
        self.kernel_size = kernel_size
        self.pool_size = pool_size
        # 子层初始化逻辑不变
        self.pool1 = MaxPooling2D(pool_size=pool_size)
        self.conv1 = Conv2D(filters=filters_1, kernel_size=kernel_size, padding='same', strides=1, activation='relu')
        self.pool2 = MaxPooling2D(pool_size=pool_size)
        self.conv2 = Conv2D(filters=filters_2, kernel_size=kernel_size, padding='same', strides=1, activation='relu')
        self.pool3 = MaxPooling2D(pool_size=pool_size)
        self.flatten = Flatten()
        self.fc1 = Dense(fc_units, activation='relu')
        self.fc2 = Dense(6, activation=None, bias_initializer=tf.keras.initializers.constant([1.0, 0.0, 0.0, 0.0, 1.0, 0.0]), kernel_initializer='zeros')
        super(Localization, self).__init__(**kwargs)

    def build(self, input_shape):
        print("Building Localization Network with input shape:", input_shape)

    def compute_output_shape(self, input_shape):
        return [None, 6]

    def call(self, inputs):
        x = self.pool1(inputs)
        x = self.conv1(x)
        x = self.pool2(x)
        x = self.conv2(x)
        x = self.pool3(x)
        x = self.flatten(x)
        x = self.fc1(x)
        theta = self.fc2(x)
        theta = tf.keras.layers.Reshape((2, 3))(theta)
        return theta

    def get_config(self):
        config = super(Localization, self).get_config()
        # 仅保留初始化需要的超参数,删除所有子层实例的配置项
        config.update({
            'filters_1': self.filters_1,
            'filters_2': self.filters_2,
            'fc_units': self.fc_units,
            'kernel_size': self.kernel_size,
            'pool_size': self.pool_size,
        })
        return config

验证加载

修复后使用你原来的加载代码即可正常加载模型:

model = load_model(BEST_MODEL_PATH, compile=False, custom_objects={'Localization': Localization})

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 20:39:00