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

如何用Fastai/PyTorch获取多分类类级Precision-Recall指标(类似ClassificationReport)?

在Fastai中获取类别级Precision-Recall指标

Fastai可结合sklearn的classification_report快速实现需求,步骤如下:

  1. 获取模型预测与真实标签
    用Fastai Learner的get_preds()方法提取验证集上的预测结果和真实标签,该方法返回两个张量:第一个是模型输出的logits,第二个是真实标签。

  2. 转换格式并计算指标
    将预测张量转换为类别索引,再和真实标签一同传入classification_report,同时传入类别名称对应关系,方便阅读结果。

示例代码:

from sklearn.metrics import classification_report

# 假设已训练好Fastai的Learner对象
preds, targs = learner.get_preds(ds_idx=1)  # ds_idx=1对应验证集,ds_idx=0对应训练集

# 将logits转换为预测类别索引
pred_classes = preds.argmax(dim=1).numpy()
# 将真实标签转换为numpy数组
true_classes = targs.numpy()

# 生成分类报告,包含每个类别的Precision、Recall、F1和样本数量
print(classification_report(
    true_classes, 
    pred_classes, 
    target_names=learner.dls.vocab  # 传入数据集的类别名称列表
))

报告中的support列会显示每个类别的样本数量,可结合Precision/Recall指标判断需要合并的小类别。


在PyTorch中获取类别级Precision-Recall指标

直接使用PyTorch时,需手动遍历验证集收集数据,再用sklearn计算指标:

  1. 设置模型为评估模式
    关闭dropout、batch norm等训练模式下的操作,避免干扰评估结果。

  2. 遍历验证集收集数据
    在无梯度环境下遍历验证集,收集每个batch的预测类别和真实标签。

  3. 计算并输出分类报告
    将收集到的数据转换为numpy数组,传入classification_report生成结果。

示例代码:

import torch
from sklearn.metrics import classification_report

# 假设已定义好模型、验证数据加载器val_loader,以及类别名称列表class_names
model.eval()
preds_list = []
targets_list = []

with torch.no_grad():
    for inputs, targets in val_loader:
        outputs = model(inputs)
        # 获取预测类别索引
        _, preds = torch.max(outputs, 1)
        # 将结果添加到列表中
        preds_list.extend(preds.cpu().numpy())
        targets_list.extend(targets.cpu().numpy())

# 生成分类报告
print(classification_report(
    targets_list, 
    preds_list, 
    target_names=class_names
))

通过support列识别样本量少的类别,结合Precision和Recall指标确定合并策略。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 05:50:35