scikit-learn GridSearchCV结合LightGBM多分类器报错排查
问题描述
- 任务目标:结合LightGBM的
gbdt算法与scikit-learn的GridSearchCV,为多分类任务筛选可靠超参数组合 - 数据情况:特征为约4000行×40列的连续值矩阵,标签为4个互斥的分类类别
- 评估方案:最初计划使用LightGBM自带的
auc_mu指标评估交叉验证折表现,当前暂选用balanced_accuracy(平衡准确率)作为评估指标
复现代码
网格搜索初始化代码
param_set = { 'n_estimators':[15, 25] } clf = lgb.LGBMModel( boosting_type='gbdt', num_leaves=31, max_depth=5, learning_rate=0.1, n_estimators=100, objective='multiclass', num_class= len(np.unique(training_data.label)), min_split_gain=0, min_child_weight=1e-3, min_child_samples=10, subsample=1, subsample_freq=0, colsample_bytree=0.6, reg_alpha=0.3, reg_lambda=0.7, random_state=42, n_jobs=2) gsearch = GridSearchCV(estimator = clf, param_grid = param_set, scoring="balanced_accuracy", error_score='raise', n_jobs=2, cv=5, verbose = 2)
拟合调用代码
# 拆分训练/验证集与测试集 stratifiedss = StratifiedShuffleSplit( n_splits = 1, test_size = 0.2, train_size = 0.8, random_state=723) for train_ind, test_ind in stratifiedss.split(X,y): train_feature_obs = X.loc[train_ind] train_labels = y[train_ind] validation_feature_obs = X.loc[test_ind] validation_labels = y[test_ind] # 转换为LightGBM Dataset格式 training_data = lgb.Dataset(train_feature_obs, label=train_labels) # 调用GridSearchCV.fit lgb_model2 = gsearch.fit(training_data.data.reset_index(drop=True), training_data.label)
报错信息
ValueError: Classification metrics can't handle a mix of unknown and continuous-multioutput targets
报错原因
报错中提到的unknown类型,来源于scikit-learn的指标校验逻辑无法识别模型输出的预测值格式:
- 你当前使用的
lgb.LGBMModel是LightGBM的底层通用估算器,没有遵循scikit-learn分类器的接口约定,其predict()方法默认返回形状为(样本数, 类别数)的类别概率矩阵(即你观测到的4个类别概率和为1的连续值输出),属于连续多输出格式 balanced_accuracy属于分类指标,要求传入的预测结果是离散的类别标签,无法直接处理连续概率矩阵,因此将无法识别的预测格式标记为unknown类型,最终抛出类型不匹配错误
解决方案
替换估算器类为适配sklearn接口的分类器
将代码中的lgb.LGBMModel替换为lgb.LGBMClassifier,该类原生遵循scikit-learn分类器API规范:- 调用
predict()方法时直接返回离散类别标签,可直接适配balanced_accuracy等需要标签输入的分类指标 - 调用
predict_proba()方法时返回类别概率矩阵,可适配AUC等需要概率输入的评估指标
修正后的分类器初始化代码如下:
clf = lgb.LGBMClassifier( boosting_type='gbdt', num_leaves=31, max_depth=5, learning_rate=0.1, n_estimators=100, objective='multiclass', min_child_weight=1e-3, min_child_samples=10, subsample=1, subsample_freq=0, colsample_bytree=0.6, reg_alpha=0.3, reg_lambda=0.7, random_state=42, n_jobs=2)注:
LGBMClassifier会自动根据输入标签识别分类类别数,无需手动传入num_class参数,手动传入也不会引发错误。- 调用
可选优化:移除多余的数据转换步骤
不需要提前将特征和标签转换为lgb.Dataset格式再传入GridSearchCV,直接传入pandas DataFrame、numpy array格式的特征和标签即可,LGBMClassifier内部会自动完成数据格式适配,可直接将拟合代码简化为:lgb_model2 = gsearch.fit(train_feature_obs.reset_index(drop=True), train_labels)如需使用概率类评估指标(如
auc_mu、多分类AUC)
可自定义scoring函数,在函数内部调用估算器的predict_proba()方法获取概率输出后完成指标计算,再将自定义函数传入GridSearchCV的scoring参数即可。
内容的提问来源于stack exchange,提问作者Rfunghifinder
相关产品推荐
相关产品推荐

