文章目录
- 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=None或ClasswiseWrapper会为那些“未被预测”的类别给出指标值(如 F1 分数为 0),从而暴露模型的短板。 - 内存友好:使用
ClasswiseWrapper时,内部实际上只维护了一个普通的多分类指标对象,因此不会增加额外的显存占用。 - 自定义标签顺序:
labels参数不仅用于命名,还决定了输出字典中键的顺序。如果传入的标签列表长度与类别数不一致,会报错,请务必保证匹配。
无论使用哪种方法,都能轻松获得每个类别上的详细指标,从而更精细地评估多分类模型的优缺点。
2、torchmetrics- ClasswiseWrapper 详解
🎯 ClasswiseWrapper:是什么?有什么用?
ClasswiseWrapper是torchmetrics中的一个包装器(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(无额外后缀))这是最新版本的完整签名,相比早期版本增加了prefix和postfix参数,提供了更强的命名可定制性。
📚 参数详解
| 参数 | 类型 | 必填 | 默认值 | 说明 |
|---|---|---|---|---|
metric | Metric | ✅ | – | 被包装的基础指标,必须是已经配置为average=None的分类指标(如MulticlassAccuracy,MulticlassF1Score等)。它内部会输出一个形状为(num_classes,)的张量。 |
labels | Optional[List[str]] | 可选 | None | 自定义的类别名称列表,长度必须与metric的类别数一致。若为None,则自动使用数字索引[0, 1, 2, ...]作为键名后缀。 |
prefix | Optional[str] | 可选 | None | 为每个输出键统一添加的前缀字符串。仅在未提供labels时生效,会替换默认的类名前缀,生成如prefix+数字的键。 |
postfix | Optional[str] | 可选 | None | 为每个输出键统一添加的后缀字符串。同样仅在未提供labels时生效,生成如数字+postfix的键。 |
关键规则:labels具有最高优先级。一旦提供了labels,prefix和postfix将被忽略,输出键固定为基础类名_标签名。
📥 输入是什么?
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的组合,键名会遵循以下层次:
- 默认行为(无
labels,无prefix/postfix)
使用基础指标的小写类名作为前缀,类别索引作为后缀,中间用下划线连接。
wrapped=ClasswiseWrapper(MulticlassAccuracy(num_classes=3,average=None))# 输出键: 'multiclassaccuracy_0', 'multiclassaccuracy_1', 'multiclassaccuracy_2'- 无
labels,但使用了prefix或postfix
此时数字索引键会直接使用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'- 提供了
labels(无论是否带prefix/postfix)
此时键名以labels为准,格式固定为基础类名_标签名。prefix和postfix会被忽略。
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内部只维护了一个基础指标实例,不增加额外的显存或计算开销。