LightGBM原生训练接口无predict_proba方法,如何输出分类概率?
LightGBM原生接口无predict_proba方法的解决方案
问题原因
LightGBM的两套使用接口设计逻辑不同:
- Scikit-learn风格的封装接口(
LGBMClassifier、LGBMRegressor等)为了对齐sklearn生态的API规范,将输出分类标签、输出分类概率的能力拆分为predict()、predict_proba()两个独立方法,符合sklearn用户的常规使用习惯。 - 原生
lgb.train()接口返回的是Booster类对象,设计上做了API简化,原生predict()方法默认直接输出分类任务的概率结果,没有单独拆分predict_proba方法,无需额外调用独立接口获取概率。
分类概率获取方法
针对你使用的二分类场景,直接调用原生模型的predict()方法即可得到正类概率:
# 接你的原生训练代码 pred_proba = model.predict(values) # 输出默认是样本对应正类(标签为1)的概率,shape为 (n_samples,) # 如果需要和sklearn接口predict_proba格式完全对齐的 [负类概率, 正类概率] 二维数组,可手动拼接: import numpy as np pred_proba_full = np.vstack([1 - pred_proba, pred_proba]).T
如果是多分类任务,原生predict()默认直接输出每个类别的概率,shape为(n_samples, n_classes),和sklearn接口predict_proba的输出格式完全一致,不需要额外处理。
注意事项
- 调用原生
predict()时如果传入了raw_score=True参数,输出的会是未经过sigmoid(二分类)/softmax(多分类)转换的原始得分,不是概率值,需要概率的话不要添加该参数即可。 - 原生
predict()输出的不是分类标签,如果需要得到0/1(二分类)或类别编号(多分类)标签,自行根据业务阈值对概率做截断即可,比如常用0.5阈值的二分类标签获取:
pred_label = (pred_proba >= 0.5).astype(int)
内容的提问来源于stack exchange,提问作者desmond
相关产品推荐
相关产品推荐

