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

如何将yellowbrick用于非Scikit模型的输出结果分析

使用Yellowbrick ClassificationReport分析PyTorch多分类模型结果的操作方法

前置准备

你需要先确保已经安装对应依赖,可通过以下命令安装:
pip install yellowbrick torch numpy matplotlib

核心逻辑说明

Yellowbrick的ClassificationReport不需要绑定训练好的PyTorch模型,只需传入测试集的真实类别标签和模型预测的类别标签即可生成可视化报告,包含每个类别的精确率、召回率、F1值和样本量。

具体操作步骤

  • 第一步:从PyTorch测试流程中获取真实标签和预测结果
    切换模型到评估模式,遍历测试集得到所有样本的真实标签和预测类别,注意要把PyTorch张量转为numpy数组,不要传入one-hot编码或者预测概率值。
  • 第二步:导入依赖并初始化可视化对象
    按你的类别索引顺序定义类别名称列表,初始化ClassificationReport实例,可通过support参数控制是否显示每个类别的样本数量,cmap参数调整图表配色。
  • 第三步:生成并输出报告
    调用score方法传入真实标签和预测标签,再调用show方法显示图表,需要保存图表时可传入outpath参数指定保存路径。

完整可运行示例代码

import torch
import numpy as np
from yellowbrick.classifier import ClassificationReport
import matplotlib.pyplot as plt

# 替换为你自己的测试集DataLoader和训练完成的模型
test_dataloader = 你的测试集DataLoader
model = 你训练好的PyTorch模型
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)

# 1、获取所有测试样本的真实标签和预测结果
model.eval()
y_true = []
y_pred = []

with torch.no_grad():
    for inputs, labels in test_dataloader:
        inputs = inputs.to(device)
        labels = labels.to(device)
        outputs = model(inputs)
        # 取置信度最高的类别作为预测结果
        _, preds = torch.max(outputs, 1)
        y_true.extend(labels.cpu().numpy())
        y_pred.extend(preds.cpu().numpy())

y_true = np.array(y_true)
y_pred = np.array(y_pred)

# 2、生成分类报告可视化
# 替换为你自己的类别名称,顺序和类别索引一一对应
classes = ["类别1", "类别2", "类别3"]
visualizer = ClassificationReport(classes=classes, support=True, cmap="Blues")

# 传入标签生成报告
visualizer.score(y_true, y_pred)

# 显示图表,需保存则添加outpath参数,如 outpath="分类报告.png"
visualizer.show()
plt.close()

常见注意事项

  • 输入的y_true和y_pred必须是形状为(样本数,)的1维数组,若存在多余维度可使用np.squeeze()处理。
  • 类别名称列表的顺序必须和训练时的类别索引完全对应,否则图表标签会匹配错误。
  • 如果模型输出的是预测概率,需要先取最大值对应的索引作为预测类别,不能直接传入概率值。

内容的提问来源于stack exchange,提问作者Dave

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 21:42:00