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

sklearn GridSearchCV与scikeras KerasClassifier配合报错求助

解决GridSearchCV搭配KerasClassifier的参数错误问题

错误原因

ValueError: Invalid parameter activation for estimator KerasClassifier 是因为你在param_grid中直接使用了activation参数,但该参数不属于KerasClassifier本身,而是属于你自定义的AnnModel构建函数的参数,GridSearchCV无法识别未关联到估算器的参数。

解决方案

步骤1:确保自定义模型函数接受目标参数

首先检查你的AnnModel函数是否定义了activation参数,示例如下:

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense

def AnnModel(activation='relu'):
    # 构建你的神经网络
    model = Sequential()
    model.add(Dense(units=32, activation=activation, input_shape=(x_train_s.shape[1],)))
    model.add(Dense(units=1, activation='sigmoid'))  # 二分类示例
    model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
    return model

步骤2:通过正确方式传递参数到GridSearchCV

有两种可行的参数传递方式:

方式一:使用build_fn__前缀指定参数

Scikeras允许通过build_fn__参数名的格式,将参数传递给build_fn指定的模型构建函数:

from scikeras.wrappers import KerasClassifier
from sklearn.model_selection import GridSearchCV

# 定义参数网格,用build_fn__前缀标记传递给AnnModel的参数
param_grid = {
    'build_fn__activation': ['relu', 'tanh', 'sigmoid'],
    # 可添加其他模型参数,比如build_fn__units: [32, 64, 128]
}

keras_wr_class = KerasClassifier(build_fn=AnnModel)
grid_search = GridSearchCV(estimator=keras_wr_class, param_grid=param_grid, scoring='accuracy', cv=5)
grid_search.fit(x_train_s, y_train)

方式二:将模型参数作为KerasClassifier构造参数传入

将AnnModel的参数直接作为KerasClassifier的构造参数,此时param_grid可直接使用参数名:

keras_wr_class = KerasClassifier(build_fn=AnnModel, activation='relu')

# 参数网格直接使用参数名
param_grid = {
    'activation': ['relu', 'tanh', 'sigmoid']
}

grid_search = GridSearchCV(estimator=keras_wr_class, param_grid=param_grid, scoring='accuracy', cv=5)
grid_search.fit(x_train_s, y_train)

验证可用参数

你可以通过以下代码查看KerasClassifier的所有可调参数,确认参数命名是否正确:

keras_wr_class = KerasClassifier(build_fn=AnnModel)
print(keras_wr_class.get_params().keys())

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 10:34:50