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

GridSearchCV调参报错:无法克隆Keras Functional模型问题咨询

问题原因与解决方案

错误原因

GridSearchCV是scikit-learn的参数调优工具,要求传入的estimator必须是符合sklearn接口的估计器——即必须实现get_params()和set_params()方法。你直接传入的Keras Functional Model对象本身不具备这些方法,因此在克隆模型时触发报错。此外,你要调整的learning_rate是优化器参数,并非模型结构参数,直接通过当前代码传递也无法生效。

解决步骤与修改后代码

1. 调整模型类,支持接收learning_rate并编译模型

在原模型类中新增方法,将模型构建与编译逻辑整合,允许传入learning_rate参数配置优化器:

from keras import backend as K, regularizers
from keras.engine.training import Model
from keras.layers import Conv2D, MaxPooling2D, Dropout, Flatten, Dense, \
BatchNormalization, Activation, Input
from keras.optimizers import Adam
import ModelLib


class Cifar100_Model(ModelLib.ModelLib):
    def build_classifier_model(self, dataset, n_classes=5,
                               activation='elu', dropout_1_rate=0.25,
                               dropout_2_rate=0.5,
                               reg_factor=50e-4, bias_reg_factor=None, batch_norm=False):
         
        n_classes = dataset.n_classes
        print(n_classes)
        print("----------------------------------------------------------------------------")
        l2_reg = regularizers.l2(reg_factor)
        l2_bias_reg = None
        if bias_reg_factor:
            l2_bias_reg = regularizers.l2(bias_reg_factor)

        # input image dimensions
        h, w, d = 32, 32, 3

        if K.image_data_format() == 'channels_first':
            input_shape = (d, h, w)
        else:
            input_shape = (h, w, d)

        x = input_1 = Input(shape=input_shape)

        x = Conv2D(filters=32, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = Conv2D(filters=32, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = MaxPooling2D(pool_size=(2, 2))(x)
        x = Dropout(rate=dropout_1_rate)(x)

        x = Conv2D(filters=64, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = Conv2D(filters=64, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = MaxPooling2D(pool_size=(2, 2))(x)
        x = Dropout(rate=dropout_1_rate)(x)

        x = Conv2D(filters=128, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = Conv2D(filters=128, kernel_size=(3, 3), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = MaxPooling2D(pool_size=(2, 2))(x)
        x = Dropout(rate=dropout_1_rate)(x)

        x = Conv2D(filters=256, kernel_size=(2, 2), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = Conv2D(filters=256, kernel_size=(2, 2), padding='same', kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)
        x = MaxPooling2D(pool_size=(2, 2))(x)
        x = Dropout(rate=dropout_1_rate)(x)

        x = Flatten()(x)
        x = Dense(units=512, kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation=activation)(x)

        x = Dropout(rate=dropout_2_rate)(x)
        x = Dense(units=n_classes, kernel_regularizer=l2_reg, bias_regularizer=l2_bias_reg)(x)
        if batch_norm:
            x = BatchNormalization()(x)
        x = Activation(activation='softmax')(x)

        model = Model(inputs=[input_1], outputs=[x])
        return model
    
    # 新增方法:构建并编译模型,接收learning_rate参数
    def create_compiled_model(self, dataset, learning_rate=0.001, **kwargs):
        model = self.build_classifier_model(dataset, **kwargs)
        # 用Adam优化器,可根据需求替换为SGD等其他优化器
        optimizer = Adam(learning_rate=learning_rate)
        # 若y_train是one-hot编码,损失函数改为categorical_crossentropy
        model.compile(optimizer=optimizer, loss='sparse_categorical_crossentropy', metrics=['accuracy'])
        return model

2. 修改测试代码,用KerasClassifier包装模型

使用KerasClassifier将Keras模型包装为sklearn兼容的估计器,再传入GridSearchCV:

from sklearn.model_selection import GridSearchCV
from keras.wrappers.scikit_learn import KerasClassifier
from functools import partial
import models.cifar100_model

def load_model():
    return models.cifar100_model.Cifar100_Model()

model_lib = load_model()

# 用partial固定dataset参数,避免GridSearchCV调用时重复传递
model_build_func = partial(model_lib.create_compiled_model, dataset=dataset)

# 包装为sklearn兼容的估计器,可在此固定训练轮数、批次大小等参数
estimator = KerasClassifier(build_fn=model_build_func, epochs=10, batch_size=32, verbose=1)

# 参数网格:key需与create_compiled_model的参数名一致
learning_rate_candidates = [0.001, 0.01, 0.1]
param_grid = dict(learning_rate=learning_rate_candidates)

# 初始化并执行网格搜索
grid = GridSearchCV(estimator=estimator, param_grid=param_grid, n_jobs=-1, cv=3, scoring='accuracy')
grid_result = grid.fit(dataset.x_train, dataset.y_train_labels)

# 输出结果
print(f"最佳验证准确率: {grid_result.best_score_:.4f}")
print(f"最佳学习率参数: {grid_result.best_params_}")

关键说明

  • KerasClassifier自动为Keras模型实现sklearn估计器所需的get_params()和set_params()方法,解决模型克隆报错问题。
  • 优化器的learning_rate需作为模型编译函数的参数传入,才能被GridSearchCV遍历调优。
  • 损失函数需根据标签格式调整:整数标签用sparse_categorical_crossentropy,one-hot编码标签用categorical_crossentropy。

内容的提问来源于stack exchange,提问作者Gustavo Henrique Nunes

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 16:15:33