如何向sklearn OneVsRestClassifier传递fit_params?遇eval_set参数错误
解决OneVsRestClassifier传递fit参数给子模型的问题
报错原因
OneVsRestClassifier的fit方法本身不接受eval_set、callbacks这类参数,这些参数属于内部包裹的LGBMClassifier的fit方法,直接传递会触发参数不匹配的报错。
解决方法
方法1:升级scikit-learn(推荐)
如果你的scikit-learn版本在0.22及以上,OneVsRestClassifier.fit()支持通过**关键字参数直接传递子模型的fit参数。确认版本后,保持原有代码结构即可;若仍报错,优先升级scikit-learn:
pip install --upgrade scikit-learn
方法2:使用fit_params参数(兼容旧版本)
针对scikit-learn 0.22以下的版本,OneVsRestClassifier.fit()提供了fit_params参数,专门用来传递给子estimator的fit方法,修改调用代码如下:
clf.fit(X_train, y_train, fit_params=fit_params)
方法3:手动训练单类别模型(备选)
如果上述方法都无法生效,可以手动遍历每个类别,单独训练LGBM模型,完全控制fit参数:
import numpy as np # 假设y_train是多标签/多分类的二维数组或稀疏矩阵 n_classes = y_train.shape[1] estimators = [] for class_idx in range(n_classes): # 针对当前类别训练模型 lgb_clf = lightgbm.LGBMClassifier(**params) lgb_clf.fit(X_train, y_train[:, class_idx], **fit_params) estimators.append(lgb_clf) # 预测示例 def predict(X): predictions = [] for clf in estimators: predictions.append(clf.predict_proba(X)[:, 1]) return np.array(predictions).T
内容的提问来源于stack exchange,提问作者Ivan Plotnikov
相关产品推荐
相关产品推荐

