如何用Fastai/PyTorch获取多分类类级Precision-Recall指标(类似ClassificationReport)?
在Fastai中获取类别级Precision-Recall指标
Fastai可结合sklearn的classification_report快速实现需求,步骤如下:
获取模型预测与真实标签
用Fastai Learner的get_preds()方法提取验证集上的预测结果和真实标签,该方法返回两个张量:第一个是模型输出的logits,第二个是真实标签。转换格式并计算指标
将预测张量转换为类别索引,再和真实标签一同传入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计算指标:
设置模型为评估模式
关闭dropout、batch norm等训练模式下的操作,避免干扰评估结果。遍历验证集收集数据
在无梯度环境下遍历验证集,收集每个batch的预测类别和真实标签。计算并输出分类报告
将收集到的数据转换为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
相关产品推荐
相关产品推荐

