You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 20:04:59