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

使用GridSearchCV遇KerasClassifier导入错误求解决方案(TF2.16.1/Keras3.3.3)

解决GridSearchCV搭配Keras/Scikeras的导入错误问题

错误原因

scikeras与当前安装的scikit-learn版本不兼容,_deprecate_Xt_in_inverse_transform函数在新版scikit-learn中已被移除,但旧版scikeras仍在调用该函数,从而引发导入错误。

具体解决办法

方法1:匹配兼容的scikeras与scikit-learn版本

  • 若使用scikit-learn 1.4.x及以上版本,安装scikeras 0.12.0及以上版本(该版本已适配新版sklearn的API变更):
    pip install --upgrade scikeras>=0.12.0
    
  • 若需保留旧版scikit-learn,可安装scikeras 0.11.0版本:
    pip install scikeras==0.11.0
    

方法2:使用Keras 3原生的scikit-learn兼容工具

Keras 3已内置对scikit-learn的支持,无需依赖第三方库,直接用以下方式导入:

import keras
from keras.wrappers.scikit_learn import KerasClassifier

方法3:降级TensorFlow/Keras版本(仅在必要时使用)

如果上述方法均不适用,可降级到TensorFlow 2.15.x和Keras 2.15.x版本,该版本的keras.wrappers.scikit_learn可正常使用:

pip install tensorflow==2.15.0 keras==2.15.0

验证步骤

安装完成后,运行以下代码测试导入是否正常:

# 方法1验证
from scikeras.wrappers import KerasClassifier
# 方法2验证
from keras.wrappers.scikit_learn import KerasClassifier

验证通过后,即可正常结合GridSearchCV使用,示例代码如下:

from sklearn.model_selection import GridSearchCV
import tensorflow as tf

def create_model(optimizer='adam'):
    model = tf.keras.models.Sequential([
        tf.keras.layers.Dense(16, activation='relu', input_shape=(8,)),
        tf.keras.layers.Dense(1, activation='sigmoid')
    ])
    model.compile(loss='binary_crossentropy', optimizer=optimizer, metrics=['accuracy'])
    return model

model = KerasClassifier(build_fn=create_model, verbose=0)
param_grid = {'optimizer': ['adam', 'sgd']}
grid = GridSearchCV(estimator=model, param_grid=param_grid, cv=3)
grid_result = grid.fit(X_train, y_train)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 08:03:15