news 2026/7/22 3:30:45

【TorchMetrics精通系列③】模型评估进阶:ClasswiseWrapper 实战与多分类细粒度指标深度解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【TorchMetrics精通系列③】模型评估进阶:ClasswiseWrapper 实战与多分类细粒度指标深度解析

文章目录

  • 1、`torchmetrics` 看到每个类别具体的指标
  • 2、`torchmetrics` - ClasswiseWrapper 详解

torchmetrics堪称模型评估界的“绝世秘籍”,招式精妙且威力无穷。若想真正参透其中玄机、融会贯通,列位看官莫急,且听我细细拆解。
这是 torchmetrics 系列文章的第三篇。
第一篇看此处:【TorchMetrics精通系列①】核心设计哲学 + Accuracy 超详解
第二篇看此处:【TorchMetrics精通系列②】混淆矩阵:归一化陷阱、TP/FP推导与10分类文本热力图分析

1、torchmetrics看到每个类别具体的指标

torchmetrics里要查看每个类别上的不同指标,主要有两种方法:

  • 直接使用指标的average=None参数
  • 或是使用ClasswiseWrapper包装器。另外,混淆矩阵也可以看作是所有类别指标的一个总览表。

方法一:使用average=None参数

这是最直接的方法。许多分类指标(如Accuracy,Precision,Recall,F1Score等)都接受一个average参数。将其设置为None'none'compute()方法就会返回一个张量,其中每个元素对应一个类别的指标值,而不是所有类别的平均值。

importtorchfromtorchmetrics.classificationimportMulticlassPrecision,MulticlassRecall# 假设是一个3分类任务num_classes=3preds=torch.randn(10,num_classes).softmax(dim=-1)# 概率target=torch.randint(num_classes,(10,))# 初始化指标时,设置 average=Noneprecision_metric=MulticlassPrecision(num_classes=num_classes,average=None)recall_metric=MulticlassRecall(num_classes=num_classes,average=None)# 累积数据precision_metric.update(preds,target)recall_metric.update(preds,target)# 获取每个类别的指标值precision_per_class=precision_metric.compute()# 形状: (num_classes,)recall_per_class=recall_metric.compute()# 形状: (num_classes,)print("每个类别的 Precision:",precision_per_class)print("每个类别的 Recall:",recall_per_class)

方法二:使用ClasswiseWrapper包装器

这是一个更灵活和强大的方法,尤其当你需要在MetricCollection中组合多个指标时。它可以将一个返回多值张量的指标(即设置了average=None的指标)“拆分”为一个字典,键名自动包含你指定的标签名,使得结果非常清晰易读。通过labels参数可以自定义类别名称。

importtorchfromtorchmetrics.wrappersimportClasswiseWrapperfromtorchmetrics.classificationimportMulticlassAccuracy num_classes=3class_names=["cat","dog","bird"]# 使用 ClasswiseWrapper 包装一个 average=None 的指标metric=ClasswiseWrapper(MulticlassAccuracy(num_classes=num_classes,average=None),labels=class_names# 给每个类别命名)# 模拟数据preds=torch.randn(10,num_classes).softmax(dim=-1)target=torch.randint(num_classes,(10,))# 计算指标,直接得到字典result=metric(preds,target)print(result)# 输出示例: {'MulticlassAccuracy_cat': tensor(0.33), 'MulticlassAccuracy_dog': tensor(0.50), 'MulticlassAccuracy_bird': tensor(0.25)}

MetricCollection结合的正确方式

将每个指标分别用ClasswiseWrapper包装,然后放入一个MetricCollection中,即可一次性获得所有指标的所有类别结果。

fromtorchmetricsimportMetricCollectionfromtorchmetrics.classificationimportMulticlassPrecision,MulticlassRecall,MulticlassF1Score num_classes=3class_names=["cat","dog","bird"]# 定义基础指标(都要设置 average=None)metrics={'Precision':MulticlassPrecision(num_classes=num_classes,average=None),'Recall':MulticlassRecall(num_classes=num_classes,average=None),'F1':MulticlassF1Score(num_classes=num_classes,average=None)}# 对每个指标使用 ClasswiseWrapper 包装,再组合成 MetricCollectionwrapped_metrics=MetricCollection({name:ClasswiseWrapper(metric_fn,labels=class_names)forname,metric_fninmetrics.items()})# 更新数据preds=torch.randn(32,num_classes).softmax(dim=-1)target=torch.randint(num_classes,(32,))wrapped_metrics.update(preds,target)# 计算所有分类别指标results=wrapped_metrics.compute()print(results)# 输出示例:# {# 'Precision_cat': tensor(0.33),# 'Precision_dog': tensor(0.50),# 'Precision_bird': tensor(0.25),# 'Recall_cat': ...,# ...# }

注意ClasswiseWrapper会自动在返回的键中加上原指标类名前缀(如Precision_cat),所以你不需要手动构造'{name}_{cls}'这样的字符串;每个指标仅需包装一次即可。


补充说明

  • 数据完整性:某些类别在数据中可能真实出现,但从未被模型预测到。这种情况下,average=NoneClasswiseWrapper会为那些“未被预测”的类别给出指标值(如 F1 分数为 0),从而暴露模型的短板。
  • 内存友好:使用ClasswiseWrapper时,内部实际上只维护了一个普通的多分类指标对象,因此不会增加额外的显存占用。
  • 自定义标签顺序labels参数不仅用于命名,还决定了输出字典中键的顺序。如果传入的标签列表长度与类别数不一致,会报错,请务必保证匹配。

无论使用哪种方法,都能轻松获得每个类别上的详细指标,从而更精细地评估多分类模型的优缺点。


2、torchmetrics- ClasswiseWrapper 详解

🎯 ClasswiseWrapper:是什么?有什么用?

ClasswiseWrappertorchmetrics中的一个包装器(wrapper),它的核心作用是将返回多值张量的分类指标(即设置了average=None的指标,每个类别一个值)“拆分”为一个更直观的字典,其中键会自动包含类别索引或你自定义的标签名。

ClasswiseWrapper的设计目标就是“透明包装”。这意味着除了最终compute()返回的结果格式变了,其他所有的使用方式都和被包装的原始指标完全一样

它解决什么问题?

当你调用F1Score(num_classes=10, average=None)时,compute()返回的是一个形状为(10,)的张量:

# 一堆数字,可读性极差tensor([0.45,0.71,0.62,0.83,0.60,0.78,0.79,0.82,0.80,0.79])

ClasswiseWrapper将它转化为:

{'f1_class_0':0.45,'f1_class_1':0.71,...}

这样在日志、TensorBoard 或控制台中都能一眼看出哪个类别表现好坏,而不需要手动对应索引。它就像一个翻译器,把“索引→值”的张量翻译成“名称→值”的字典。


📝 完整函数签名

classtorchmetrics.wrappers.ClasswiseWrapper(metric:Metric,# 必填:被包装的基础指标labels:Optional[List[str]]=None,# 可选,默认 None(自动使用数字索引)prefix:Optional[str]=None,# 可选,默认 None(无额外前缀)postfix:Optional[str]=None# 可选,默认 None(无额外后缀))

这是最新版本的完整签名,相比早期版本增加了prefixpostfix参数,提供了更强的命名可定制性。


📚 参数详解

参数类型必填默认值说明
metricMetric被包装的基础指标,必须是已经配置为average=None的分类指标(如MulticlassAccuracy,MulticlassF1Score等)。它内部会输出一个形状为(num_classes,)的张量。
labelsOptional[List[str]]可选None自定义的类别名称列表,长度必须与metric的类别数一致。若为None,则自动使用数字索引[0, 1, 2, ...]作为键名后缀。
prefixOptional[str]可选None为每个输出键统一添加的前缀字符串。仅在未提供labels时生效,会替换默认的类名前缀,生成如prefix+数字的键。
postfixOptional[str]可选None为每个输出键统一添加的后缀字符串。同样仅在未提供labels时生效,生成如数字+postfix的键。

关键规则labels具有最高优先级。一旦提供了labelsprefixpostfix将被忽略,输出键固定为基础类名_标签名


📥 输入是什么?

ClasswiseWrapper本身是一个包装器,它不改变底层指标的输入要求。你调用update()forward()时传入的参数,和直接使用被包装的指标时完全一样。

对于多分类任务:

情况preds形状preds类型target形状target类型
传入概率/logits(N, C)float32(N,)long
传入预测类别索引(N,)long(N,)long

📤 输出结果是什么?

  • compute():返回一个字典Dict[str, Tensor]),键名根据参数自动生成,值为标量张量(每个类别的指标值)。
  • forward(*args, **kwargs)或直接调用:等价于先update()compute(),返回同样的字典。

键名生成规则详解

根据是否提供labels以及prefix/postfix的组合,键名会遵循以下层次:

  1. 默认行为(无labels,无prefix/postfix

使用基础指标的小写类名作为前缀,类别索引作为后缀,中间用下划线连接。

wrapped=ClasswiseWrapper(MulticlassAccuracy(num_classes=3,average=None))# 输出键: 'multiclassaccuracy_0', 'multiclassaccuracy_1', 'multiclassaccuracy_2'
  1. labels,但使用了prefixpostfix

此时数字索引键会直接使用prefix/postfix不再包含基础指标类名

# prefix 示例:直接用前缀 + 数字ClasswiseWrapper(MulticlassAccuracy(num_classes=3,average=None),prefix="acc-")# 输出键: 'acc-0', 'acc-1', 'acc-2'# postfix 示例:数字 + 后缀ClasswiseWrapper(MulticlassAccuracy(num_classes=3,average=None),postfix="-acc")# 输出键: '0-acc', '1-acc', '2-acc'
  1. 提供了labels(无论是否带prefix/postfix

此时键名以labels为准,格式固定为基础类名_标签名prefixpostfix会被忽略。

wrapped=ClasswiseWrapper(MulticlassF1Score(num_classes=2,average=None),labels=["negative","positive"],prefix="val_"# 该 prefix 不会生效)# 输出键: 'multiclassf1score_negative', 'multiclassf1score_positive'

与早期版本的区别

如果你查阅的是旧版文档(如 v0.9.0),会发现键名可能不含完整的类名前缀:

# 旧版本(v0.9.0):# {'accuracy_0': ..., 'accuracy_horse': ...}# 新版本(v1.0+):# {'multiclassaccuracy_0': ..., 'multiclassaccuracy_horse': ...}

这是因为新版本使用基础指标的完整小写类名(如'multiclassaccuracy')而非简写(如'accuracy')作为默认前缀,避免了不同指标类型输出键名冲突的问题。


⚙️ 常用操作

① 基础使用(默认数字索引)

importtorchfromtorchmetrics.wrappersimportClasswiseWrapperfromtorchmetrics.classificationimportMulticlassAccuracy# 必须设置 average=Nonemetric=ClasswiseWrapper(MulticlassAccuracy(num_classes=10,average=None))forbatchinval_loader:preds,target=batch metric.update(preds,target)result=metric.compute()print(result)# {'multiclassaccuracy_0': tensor(0.45), ..., 'multiclassaccuracy_9': tensor(0.78)}metric.reset()

② 使用自定义标签名

class_names=["科技","体育","财经","娱乐","教育","军事","健康","农业","游戏","房产"]wrapped=ClasswiseWrapper(MulticlassF1Score(num_classes=10,average=None),labels=class_names)# 输出: {'multiclassf1score_科技': 0.45, 'multiclassf1score_体育': 0.71, ...}

③ 使用 prefix 快速区分训练/验证

# 训练阶段(无 labels,仅用 prefix)train_acc=ClasswiseWrapper(MulticlassAccuracy(num_classes=10,average=None),prefix="train_")# 输出: {'train_0': ..., 'train_1': ...}# 验证阶段val_acc=ClasswiseWrapper(MulticlassAccuracy(num_classes=10,average=None),prefix="val_")# 输出: {'val_0': ..., 'val_1': ...}

④ 在 PyTorch Lightning 中使用

classMyModel(pl.LightningModule):def__init__(self):super().__init__()self.val_metrics=MetricCollection({'f1':ClasswiseWrapper(MulticlassF1Score(num_classes=10,average=None),labels=class_names)})defvalidation_step(self,batch,batch_idx):...self.val_metrics.update(preds,target)# 可用于 log_dictself.log_dict(self.val_metrics,on_step=False,on_epoch=True)

💎 核心要点

  • 前提条件:被包装的指标必须设置average=None,使其输出逐类别的张量。
  • 键名优先级labels>prefix/postfix。若提供了labels,键名固定为“基础类名_标签名”;若未提供labels但提供了prefix/postfix,键名变为“前缀+数字”或“数字+后缀”;默认则为“基础类名_数字”。
  • 团队建议:优先使用labels自定义类别名,可读性最佳;prefix/postfix适合快速区分 train/val 阶段且不关心具体类别名的场景。
  • 无缝组合:与MetricCollection结合后,所有指标的所有类别结果会被扁平化到一个字典中,非常适合一次性记录或日志输出。
  • 零额外开销ClasswiseWrapper内部只维护了一个基础指标实例,不增加额外的显存或计算开销。


版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/22 3:30:36

抖音直播伴侣画质优化与参数配置完整方案(附各场景对照表)

前言 抖音直播伴侣是抖音官方推出的PC端直播工具,支持游戏直播、聊天互动、电商带货等多种场景。但很多人在使用中遇到画质模糊、卡顿掉帧、声音异常等问题——排查下来,大部分不是硬件不够,而是参数没配对。 这篇文章整理了一套经过实测的…

作者头像 李华
网站建设 2026/7/22 3:29:25

泰安AI数字人公司哪家靠谱?2026年本地服务商技术实力与落地能力盘点

随着人工智能技术的深度演进,2026年的企业服务与公共传播领域正迎来一场安静却深刻的改变。从展厅导览、政务导办,到文化IP活化、品牌内容创作,AI数字人正逐渐从新颖的概念转变为诸多行业提升服务效率的实用工具。 不过,面对市场上…

作者头像 李华
网站建设 2026/7/22 3:28:22

嵌入式系统引导:NAND Flash与MMC/SD卡启动原理与工程实践

1. 嵌入式系统引导:从存储介质到第一行代码的旅程当一块嵌入式芯片上电,从一片“黑暗”到执行我们编写的应用程序,这中间发生了什么?对于很多开发者来说,这个过程像一个黑盒:把编译好的镜像烧录到存储芯片&…

作者头像 李华
网站建设 2026/7/22 3:26:02

Spider RPC 2.0深度解析:性能优化与分布式通信实践

1. Spider RPC 2.0.0-RELEASE版本深度解析作为分布式系统通信的核心组件,RPC框架的每一次重大版本升级都值得开发者高度关注。Spider RPC 2.0.0-RELEASE的发布标志着该框架在性能、功能和稳定性方面都达到了新的高度。本文将带您深入剖析这次更新的技术细节&#xf…

作者头像 李华
网站建设 2026/7/22 3:25:25

很多内耗,本质是价值冲突

内耗不是因为你想得太多,而是因为你的内心存在多个方向,它们都在争夺你的行动权。很多时候: 你不是不知道该做什么。 而是: 你同时想要互相冲突的东西。第一层:什么是价值冲突? 价值冲突: 就是两…

作者头像 李华
网站建设 2026/7/22 3:24:39

Java中间件实战01:多环境Redis部署最全指南|Windows/Linux/Docker版本纠错+生产配置+SpringBoot3整合

文章目录 📌 专栏信息 一、前言(🔥痛点前置) 1.1 行业通用痛点(90%开发者必踩坑❌) 1.2 本文核心学习收益(💎干货价值) 1.3 技术版本适配说明(🚨避坑核心) 二、💻 Windows环境Redis部署(本地开发专属) 2.1 下载与环境规范(🚨避坑前置) 2.2 两种启动方式…

作者头像 李华