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

Keras使用KerasClassifier与cross_val_score报TypeError的解决方法咨询

解决KerasClassifier交叉验证时的Pickle错误

嘿,我一眼就揪出问题所在了——你在初始化KerasClassifier的时候踩了个很容易犯的小坑!

错误根源

看这行代码:

classfier = KerasClassifier(build_fn=func1(),batch_size=10, epochs=100)

你把func1()的执行结果(也就是已经建好的神经网络模型对象)传给了build_fn参数,但KerasClassifier要求build_fn是一个函数本身,而不是已经实例化好的模型。

交叉验证时scikit-learn会尝试克隆你的估算器来生成多个训练实例,但Keras模型内部包含无法被pickle序列化的_thread.lock对象,这就导致了那串冗长的报错信息。

另外,你的数据标准化环节也有小疏漏:scale.fit_transform(Xtrain)没有把结果赋值回Xtrain,相当于标准化操作白做了,这会影响模型的训练效果。

修复方案

只需要做两处关键修改:

  1. 把build_fn=func1()改成build_fn=func1(传递函数对象,而非执行后的模型)
  2. 修正标准化代码,把处理后的数据重新赋值给训练集和测试集

修正后的完整代码

import pandas as pd 
import numpy as np 
import matplotlib.pyplot as plt 
dataset = pd.read_csv("Churn_Modelling.csv") 
X = dataset.iloc[:,3:13].values 
Y = dataset.iloc[:,13:].values 

from sklearn.preprocessing import OneHotEncoder,LabelEncoder,StandardScaler 
enc1=LabelEncoder() 
enc2=LabelEncoder() 
X[:,1] = enc1.fit_transform(X[:,1]) 
X[:,2] = enc2.fit_transform(X[:,2]) 

# 注:新版本scikit-learn中OneHotEncoder的categorical_features已废弃,建议改用ColumnTransformer
one = OneHotEncoder(categorical_features=[1]) 
X=one.fit_transform(X).toarray() 
X = X[:,1:] 

from sklearn.model_selection import train_test_split 
Xtrain,Xtest,Ytrain,Ytest = train_test_split(X,Y,random_state=0,test_size=0.2) 

# 修复标准化:将处理后的数据赋值回原变量
scale = StandardScaler() 
Xtrain = scale.fit_transform(Xtrain) 
Xtest = scale.transform(Xtest) 

from keras.wrappers.scikit_learn import KerasClassifier 
from sklearn.model_selection import cross_val_score 
from keras.models import Sequential 
from keras.layers import Dense 

def func1(): 
    net = Sequential() 
    net.add(Dense(input_dim=11,units=6,activation="relu",kernel_initializer='uniform')) 
    net.add(Dense(units=6,activation="relu",kernel_initializer='uniform')) 
    net.add(Dense(units=1,activation="sigmoid",kernel_initializer='uniform')) 
    net.compile(optimizer='adam',metrics=['accuracy'],loss='binary_crossentropy') 
    return net 

# 关键修改:build_fn传入函数名func1,而非执行后的模型对象
classfier = KerasClassifier(build_fn=func1, batch_size=10, epochs=100) 
cross = cross_val_score(estimator=classfier, X=Xtrain, y=Ytrain, cv=10, n_jobs=-1)

# 可以打印交叉验证的结果,方便查看效果
print("交叉验证准确率均值:", cross.mean())
print("交叉验证准确率标准差:", cross.std())

额外提示:如果用的是新版本scikit-learn,OneHotEncoder的categorical_features参数已经被移除,建议改用ColumnTransformer来指定需要编码的列,避免后续版本兼容问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:35:51