使用sklearn的RandomizedSearchCV调参触发KerasRegressor克隆RuntimeError如何解决
问题原因
这个错误不是你参数格式写错导致的,根源是TensorFlow自带的KerasRegressor/KerasClassifier包装器设计不符合高版本Scikit-learn的克隆规则:sklearn的clone方法会校验 estimator 构造函数的参数和实例存储的参数是否一致,而旧版TF的keras包装器没有把你传入的自定义模型参数(比如learning_rate2、nn21等)正确注册为实例属性,克隆的时候检测不到对应参数就会抛出异常。
无需降级sklearn的解决方案
方案1:改用官方维护的SciKeras包装器(最推荐)
SciKeras是Keras官方现在维护的、专门用于兼容Scikit-learn的包装库,完全兼容现有调参逻辑,代码改动极小:
- 先安装依赖:
pip install scikeras - 替换导入语句:
把原来的from tensorflow.keras.wrappers.scikit_learn import KerasRegressor
替换为from scikeras.wrappers import KerasRegressor
其余代码不需要做任何修改,即可正常运行随机搜索。
方案2:关闭多进程(零代码修改临时方案)
克隆错误绝大多数是在多进程并行调参时触发的,你可以把RandomizedSearchCV的n_jobs参数从-1改为1,单进程运行不需要多次克隆estimator实例,即可绕过这个错误。缺点是调参速度会变慢,适合参数搜索空间小的场景:
grid = RandomizedSearchCV(model3, param_distributions=param_dist, n_iter=10, n_jobs=1, cv=5, scoring='neg_mean_absolute_error')
方案3:自定义兼容克隆规则的包装类
如果你不想额外安装依赖也不想用单进程,可以自己写一个继承原生KerasRegressor的类,重写get_params和set_params方法满足sklearn的校验要求:
from tensorflow.keras.wrappers.scikit_learn import KerasRegressor as BaseKerasRegressor class KerasRegressor(BaseKerasRegressor): def get_params(self, **params): res = super().get_params(**params) res.update(self.sk_params) return res def set_params(self, **params): self.sk_params = params return super().set_params(**params)
然后用你自定义的KerasRegressor实例化model3即可。
内容的提问来源于stack exchange,提问作者WieWie
相关产品推荐
相关产品推荐

