文章目录
- [1、`torchmetrics` 看到每个类别具体的指标](#1、
torchmetrics看到每个类别具体的指标) - [2、`torchmetrics` - ClasswiseWrapper 详解](#2、
torchmetrics- ClasswiseWrapper 详解)
torchmetrics堪称模型评估界的"绝世秘籍",招式精妙且威力无穷。若想真正参透其中玄机、融会贯通,列位看官莫急,且听我细细拆解。这是 torchmetrics 系列文章的第三篇。
1、torchmetrics 看到每个类别具体的指标
在 torchmetrics 里要查看每个类别上的不同指标,主要有两种方法:
- 直接使用指标的
average=None参数 - 或是使用
ClasswiseWrapper包装器。另外,混淆矩阵也可以看作是所有类别指标的一个总览表。
方法一:使用 average=None 参数
这是最直接的方法。许多分类指标(如 Accuracy, Precision, Recall, F1Score 等)都接受一个 average 参数。将其设置为 None 或 'none',compute() 方法就会返回一个张量,其中每个元素对应一个类别的指标值,而不是所有类别的平均值。
python
import torch
from torchmetrics.classification import MulticlassPrecision, MulticlassRecall
# 假设是一个3分类任务
num_classes = 3
preds = torch.randn(10, num_classes).softmax(dim=-1) # 概率
target = torch.randint(num_classes, (10,))
# 初始化指标时,设置 average=None
precision_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 参数可以自定义类别名称。
python
import torch
from torchmetrics.wrappers import ClasswiseWrapper
from torchmetrics.classification import MulticlassAccuracy
num_classes = 3
class_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 中,即可一次性获得所有指标的所有类别结果。
python
from torchmetrics import MetricCollection
from torchmetrics.classification import MulticlassPrecision, MulticlassRecall, MulticlassF1Score
num_classes = 3
class_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 包装,再组合成 MetricCollection
wrapped_metrics = MetricCollection({
name: ClasswiseWrapper(metric_fn, labels=class_names)
for name, metric_fn in metrics.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,) 的张量:
python
# 一堆数字,可读性极差
tensor([0.45, 0.71, 0.62, 0.83, 0.60, 0.78, 0.79, 0.82, 0.80, 0.79])
ClasswiseWrapper 将它转化为:
python
{
'f1_class_0': 0.45, 'f1_class_1': 0.71, ...
}
这样在日志、TensorBoard 或控制台中都能一眼看出哪个类别表现好坏,而不需要手动对应索引。它就像一个翻译器,把"索引→值"的张量翻译成"名称→值"的字典。
📝 完整函数签名
python
class torchmetrics.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)
使用基础指标的小写类名作为前缀,类别索引作为后缀,中间用下划线连接。
python
wrapped = ClasswiseWrapper(
MulticlassAccuracy(num_classes=3, average=None)
)
# 输出键: 'multiclassaccuracy_0', 'multiclassaccuracy_1', 'multiclassaccuracy_2'
- 无
labels,但使用了prefix或postfix
此时数字索引键会直接使用 prefix/postfix,不再包含基础指标类名。
python
# 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 会被忽略。
python
wrapped = ClasswiseWrapper(
MulticlassF1Score(num_classes=2, average=None),
labels=["negative", "positive"],
prefix="val_" # 该 prefix 不会生效
)
# 输出键: 'multiclassf1score_negative', 'multiclassf1score_positive'
与早期版本的区别
如果你查阅的是旧版文档(如 v0.9.0),会发现键名可能不含完整的类名前缀:
python
# 旧版本(v0.9.0):
# {'accuracy_0': ..., 'accuracy_horse': ...}
# 新版本(v1.0+):
# {'multiclassaccuracy_0': ..., 'multiclassaccuracy_horse': ...}
这是因为新版本使用基础指标的完整小写类名 (如 'multiclassaccuracy')而非简写(如 'accuracy')作为默认前缀,避免了不同指标类型输出键名冲突的问题。
⚙️ 常用操作
① 基础使用(默认数字索引)
python
import torch
from torchmetrics.wrappers import ClasswiseWrapper
from torchmetrics.classification import MulticlassAccuracy
# 必须设置 average=None
metric = ClasswiseWrapper(
MulticlassAccuracy(num_classes=10, average=None)
)
for batch in val_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()
② 使用自定义标签名
python
class_names = ["科技", "体育", "财经", "娱乐", "教育", "军事", "健康", "农业", "游戏", "房产"]
wrapped = ClasswiseWrapper(
MulticlassF1Score(num_classes=10, average=None),
labels=class_names
)
# 输出: {'multiclassf1score_科技': 0.45, 'multiclassf1score_体育': 0.71, ...}
③ 使用 prefix 快速区分训练/验证
python
# 训练阶段(无 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 中使用
python
class MyModel(pl.LightningModule):
def __init__(self):
super().__init__()
self.val_metrics = MetricCollection({
'f1': ClasswiseWrapper(
MulticlassF1Score(num_classes=10, average=None),
labels=class_names
)
})
def validation_step(self, batch, batch_idx):
...
self.val_metrics.update(preds, target)
# 可用于 log_dict
self.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内部只维护了一个基础指标实例,不增加额外的显存或计算开销。