使用KerasRegressor搭配cross_validate因不可克隆性报错如何解决
问题根源
问题出在自定义评估器不符合sklearn的实现规范,sklearn的clone()函数会通过get_params()获取评估器所有初始化参数,再传入构造函数生成新实例用于交叉验证的不同折。你之前的代码没有将自定义的batch_input_shape纳入父类参数管理,也没有正确绑定模型构建函数,导致克隆后的实例丢失必要属性,触发KeyError: 'batch_input_shape'。
修复后可直接运行的代码
from sklearn.datasets import make_regression from sklearn.model_selection import cross_validate from tensorflow.keras.layers import Dense, LSTM from tensorflow.keras.models import Sequential from tensorflow.keras.wrappers.scikit_learn import KerasRegressor class MyRegressor(KerasRegressor): def __init__(self, batch_input_shape, **kwargs): self.batch_input_shape = batch_input_shape # 显式绑定模型构建方法给父类 super().__init__(build_fn=self.__call__, **kwargs) def __call__(self): model = Sequential([ LSTM(16, stateful=True, batch_input_shape=self.batch_input_shape), Dense(1), ]) model.compile(optimizer='adam', loss='mean_squared_error', metrics=['RootMeanSquaredError']) return model def reset_states(self): if hasattr(self, 'model'): self.model.reset_states() X, y = make_regression(6400, 5) X = X.reshape(X.shape[0], 1, X.shape[1]) batch_size = 64 batch_input_shape = (batch_size, 1, X.shape[-1]) # 单实例运行验证 reg = MyRegressor(batch_input_shape, epochs=10, verbose=0) for i in range(10): reg.fit(X, y, batch_size=batch_size) reg.reset_states() # 交叉验证运行 reg = MyRegressor(batch_input_shape, epochs=10, verbose=0) results = cross_validate(reg, X, y, scoring=['neg_mean_squared_error'], fit_params={'batch_size': batch_size}) print(results)
核心修改说明
- 调整
__init__逻辑,显式将当前类的模型构建方法__call__作为build_fn参数传入父类,确保Keras包装器能正确识别模型构建逻辑 - 移除了冗余的
build_fn=None参数,避免参数冲突导致get_params()无法捕获自定义的batch_input_shape - 给
reset_states方法增加属性存在性判断,避免模型未初始化时调用报错 - 调用
cross_validate时通过fit_params显式传入训练所需的batch_size,和batch_input_shape的配置匹配 - 初始化评估器时可直接传入训练相关参数(如epochs、verbose),符合sklearn评估器的使用规范
注意:
tf.keras.wrappers.scikit_learn.KerasRegressor在TensorFlow 2.7及以上版本已被标记为废弃,若使用新版本TensorFlow建议切换到SciKeras的KerasRegressor实现,适配逻辑一致。
内容的提问来源于stack exchange,提问作者Felix
相关产品推荐
相关产品推荐

