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,相当于标准化操作白做了,这会影响模型的训练效果。
修复方案
只需要做两处关键修改:
- 把
build_fn=func1()改成build_fn=func1(传递函数对象,而非执行后的模型) - 修正标准化代码,把处理后的数据重新赋值给训练集和测试集
修正后的完整代码
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
相关产品推荐
相关产品推荐

