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

【TorchMetrics精通系列③】模型评估进阶:ClasswiseWrapper 实战与多分类细粒度指标深度解析 文章目录1、torchmetrics 看到每个类别具体的指标2、torchmetrics - ClasswiseWrapper 详解torchmetrics堪称模型评估界的“绝世秘籍”招式精妙且威力无穷。若想真正参透其中玄机、融会贯通列位看官莫急且听我细细拆解。这是 torchmetrics 系列文章的第三篇。第一篇看此处【TorchMetrics精通系列①】核心设计哲学 Accuracy 超详解第二篇看此处【TorchMetrics精通系列②】混淆矩阵归一化陷阱、TP/FP推导与10分类文本热力图分析1、torchmetrics看到每个类别具体的指标在torchmetrics里要查看每个类别上的不同指标主要有两种方法直接使用指标的averageNone参数或是使用ClasswiseWrapper包装器。另外混淆矩阵也可以看作是所有类别指标的一个总览表。方法一使用averageNone参数这是最直接的方法。许多分类指标如Accuracy,Precision,Recall,F1Score等都接受一个average参数。将其设置为None或nonecompute()方法就会返回一个张量其中每个元素对应一个类别的指标值而不是所有类别的平均值。importtorchfromtorchmetrics.classificationimportMulticlassPrecision,MulticlassRecall# 假设是一个3分类任务num_classes3predstorch.randn(10,num_classes).softmax(dim-1)# 概率targettorch.randint(num_classes,(10,))# 初始化指标时设置 averageNoneprecision_metricMulticlassPrecision(num_classesnum_classes,averageNone)recall_metricMulticlassRecall(num_classesnum_classes,averageNone)# 累积数据precision_metric.update(preds,target)recall_metric.update(preds,target)# 获取每个类别的指标值precision_per_classprecision_metric.compute()# 形状: (num_classes,)recall_per_classrecall_metric.compute()# 形状: (num_classes,)print(每个类别的 Precision:,precision_per_class)print(每个类别的 Recall:,recall_per_class)方法二使用ClasswiseWrapper包装器这是一个更灵活和强大的方法尤其当你需要在MetricCollection中组合多个指标时。它可以将一个返回多值张量的指标即设置了averageNone的指标“拆分”为一个字典键名自动包含你指定的标签名使得结果非常清晰易读。通过labels参数可以自定义类别名称。importtorchfromtorchmetrics.wrappersimportClasswiseWrapperfromtorchmetrics.classificationimportMulticlassAccuracy num_classes3class_names[cat,dog,bird]# 使用 ClasswiseWrapper 包装一个 averageNone 的指标metricClasswiseWrapper(MulticlassAccuracy(num_classesnum_classes,averageNone),labelsclass_names# 给每个类别命名)# 模拟数据predstorch.randn(10,num_classes).softmax(dim-1)targettorch.randint(num_classes,(10,))# 计算指标直接得到字典resultmetric(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_classes3class_names[cat,dog,bird]# 定义基础指标都要设置 averageNonemetrics{Precision:MulticlassPrecision(num_classesnum_classes,averageNone),Recall:MulticlassRecall(num_classesnum_classes,averageNone),F1:MulticlassF1Score(num_classesnum_classes,averageNone)}# 对每个指标使用 ClasswiseWrapper 包装再组合成 MetricCollectionwrapped_metricsMetricCollection({name:ClasswiseWrapper(metric_fn,labelsclass_names)forname,metric_fninmetrics.items()})# 更新数据predstorch.randn(32,num_classes).softmax(dim-1)targettorch.randint(num_classes,(32,))wrapped_metrics.update(preds,target)# 计算所有分类别指标resultswrapped_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}这样的字符串每个指标仅需包装一次即可。补充说明数据完整性某些类别在数据中可能真实出现但从未被模型预测到。这种情况下averageNone或ClasswiseWrapper会为那些“未被预测”的类别给出指标值如 F1 分数为 0从而暴露模型的短板。内存友好使用ClasswiseWrapper时内部实际上只维护了一个普通的多分类指标对象因此不会增加额外的显存占用。自定义标签顺序labels参数不仅用于命名还决定了输出字典中键的顺序。如果传入的标签列表长度与类别数不一致会报错请务必保证匹配。无论使用哪种方法都能轻松获得每个类别上的详细指标从而更精细地评估多分类模型的优缺点。2、torchmetrics- ClasswiseWrapper 详解 ClasswiseWrapper是什么有什么用ClasswiseWrapper是torchmetrics中的一个包装器wrapper它的核心作用是将返回多值张量的分类指标即设置了averageNone的指标每个类别一个值“拆分”为一个更直观的字典其中键会自动包含类别索引或你自定义的标签名。ClasswiseWrapper的设计目标就是“透明包装”。这意味着除了最终compute()返回的结果格式变了其他所有的使用方式都和被包装的原始指标完全一样。它解决什么问题当你调用F1Score(num_classes10, averageNone)时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参数提供了更强的命名可定制性。 参数详解参数类型必填默认值说明metricMetric✅–被包装的基础指标必须是已经配置为averageNone的分类指标如MulticlassAccuracy,MulticlassF1Score等。它内部会输出一个形状为(num_classes,)的张量。labelsOptional[List[str]]可选None自定义的类别名称列表长度必须与metric的类别数一致。若为None则自动使用数字索引[0, 1, 2, ...]作为键名后缀。prefixOptional[str]可选None为每个输出键统一添加的前缀字符串。仅在未提供labels时生效会替换默认的类名前缀生成如prefix数字的键。postfixOptional[str]可选None为每个输出键统一添加的后缀字符串。同样仅在未提供labels时生效生成如数字postfix的键。关键规则labels具有最高优先级。一旦提供了labelsprefix和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使用基础指标的小写类名作为前缀类别索引作为后缀中间用下划线连接。wrappedClasswiseWrapper(MulticlassAccuracy(num_classes3,averageNone))# 输出键: multiclassaccuracy_0, multiclassaccuracy_1, multiclassaccuracy_2无labels但使用了prefix或postfix此时数字索引键会直接使用prefix/postfix不再包含基础指标类名。# prefix 示例直接用前缀 数字ClasswiseWrapper(MulticlassAccuracy(num_classes3,averageNone),prefixacc-)# 输出键: acc-0, acc-1, acc-2# postfix 示例数字 后缀ClasswiseWrapper(MulticlassAccuracy(num_classes3,averageNone),postfix-acc)# 输出键: 0-acc, 1-acc, 2-acc提供了labels无论是否带prefix/postfix此时键名以labels为准格式固定为基础类名_标签名。prefix和postfix会被忽略。wrappedClasswiseWrapper(MulticlassF1Score(num_classes2,averageNone),labels[negative,positive],prefixval_# 该 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# 必须设置 averageNonemetricClasswiseWrapper(MulticlassAccuracy(num_classes10,averageNone))forbatchinval_loader:preds,targetbatch metric.update(preds,target)resultmetric.compute()print(result)# {multiclassaccuracy_0: tensor(0.45), ..., multiclassaccuracy_9: tensor(0.78)}metric.reset()② 使用自定义标签名class_names[科技,体育,财经,娱乐,教育,军事,健康,农业,游戏,房产]wrappedClasswiseWrapper(MulticlassF1Score(num_classes10,averageNone),labelsclass_names)# 输出: {multiclassf1score_科技: 0.45, multiclassf1score_体育: 0.71, ...}③ 使用 prefix 快速区分训练/验证# 训练阶段无 labels仅用 prefixtrain_accClasswiseWrapper(MulticlassAccuracy(num_classes10,averageNone),prefixtrain_)# 输出: {train_0: ..., train_1: ...}# 验证阶段val_accClasswiseWrapper(MulticlassAccuracy(num_classes10,averageNone),prefixval_)# 输出: {val_0: ..., val_1: ...}④ 在 PyTorch Lightning 中使用classMyModel(pl.LightningModule):def__init__(self):super().__init__()self.val_metricsMetricCollection({f1:ClasswiseWrapper(MulticlassF1Score(num_classes10,averageNone),labelsclass_names)})defvalidation_step(self,batch,batch_idx):...self.val_metrics.update(preds,target)# 可用于 log_dictself.log_dict(self.val_metrics,on_stepFalse,on_epochTrue) 核心要点前提条件被包装的指标必须设置averageNone使其输出逐类别的张量。键名优先级labelsprefix/postfix。若提供了labels键名固定为“基础类名_标签名”若未提供labels但提供了prefix/postfix键名变为“前缀数字”或“数字后缀”默认则为“基础类名_数字”。团队建议优先使用labels自定义类别名可读性最佳prefix/postfix适合快速区分 train/val 阶段且不关心具体类别名的场景。无缝组合与MetricCollection结合后所有指标的所有类别结果会被扁平化到一个字典中非常适合一次性记录或日志输出。零额外开销ClasswiseWrapper内部只维护了一个基础指标实例不增加额外的显存或计算开销。