如何用XGBClassifier获取验证集及新数据的Top3类别预测概率?
解决方案
一、先修正现有代码的笔误
你的代码里le.classed_是拼写错误,应该改为le.classes_,否则会触发AttributeError。
二、新数据的Top3类别预测
对于预处理好的新数据向量,先通过模型得到所有类别概率,再排序提取Top3,最后映射回原始类别标签:
import numpy as np # 假设new_data是已转换为向量的新输入数据 new_pred_prob = xgb_cl.predict_proba(new_data) # 对每个样本的概率数组,从高到低取Top3的索引和对应的概率值 top3_indices = np.argsort(new_pred_prob, axis=1)[:, -3:][:, ::-1] top3_probs = np.take_along_axis(new_pred_prob, top3_indices, axis=1) # 将索引映射回原始类别标签,生成每个样本的Top3结果字典 top3_preds = [] for idx_arr, prob_arr in zip(top3_indices, top3_probs): top3_classes = le.inverse_transform(idx_arr) top3_preds.append(dict(zip(top3_classes, prob_arr))) # 输出示例 for idx, res in enumerate(top3_preds): print(f"第{idx+1}条数据Top3预测:{res}")
三、测试集上的Top3概率验证
针对测试集,我们可以从两个维度做验证:一是统计真实类别是否落在Top3预测里,二是查看真实类别概率的整体排名。
方式1:统计真实类别在Top3中的命中率
# 处理测试集的概率结果 test_top3_indices = np.argsort(pred_prob, axis=1)[:, -3:][:, ::-1] test_top3_probs = np.take_along_axis(pred_prob, test_top3_indices, axis=1) test_top3_classes = le.inverse_transform(test_top3_indices.flatten()).reshape(test_top3_indices.shape) # 注意:如果y_test是原始类别标签,直接用y_test即可;如果是编码后的,用le.inverse_transform(y_test)转换 y_test_true = y_test # 生成验证结果列表 validation_list = [] for true_cls, top3_cls, top3_prob in zip(y_test_true, test_top3_classes, test_top3_probs): validation_list.append({ "真实类别": true_cls, "Top3预测(类别:概率)": dict(zip(top3_cls, top3_prob)), "真实类别是否在Top3": true_cls in top3_cls }) # 计算Top3命中率 hit_rate = sum(item["真实类别是否在Top3"] for item in validation_list) / len(validation_list) print(f"测试集Top3预测命中率:{hit_rate:.2%}")
方式2:查看真实类别概率的全局排名
# 获取每个测试样本真实类别对应的概率值 true_cls_indices = le.transform(y_test_true) true_cls_probs = pred_prob[np.arange(len(pred_prob)), true_cls_indices] # 计算真实类别概率的排名(从高到低,排名1为最高概率) true_cls_ranks = np.argsort(-pred_prob, axis=1).argsort(axis=1)[np.arange(len(pred_prob)), true_cls_indices] + 1 # 生成排名结果 rank_results = [] for true_cls, prob, rank in zip(y_test_true, true_cls_probs, true_cls_ranks): rank_results.append({ "真实类别": true_cls, "真实类别概率": round(prob, 4), "真实类别概率排名": rank }) # 输出前5条示例 print("测试集真实类别概率排名示例:") for res in rank_results[:5]: print(res)
内容的提问来源于stack exchange,提问作者Aditya sharma
相关产品推荐
相关产品推荐

