You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch Lightning的test_step中添加每类精确率、召回率、F1指标?

PyTorch Lightning test_step中每类指标计算的问题与解决方法

你代码中的错误点

  1. 指标更新逻辑错误:直接调用self.precision_each(preds, labels)只会计算当前batch的指标,无法累积整个测试集的结果。torchmetrics的指标类需要通过update方法逐步累积数据,最终用compute得到全局统计结果。
  2. 拼写错误:代码中self.logger.expe是笔误,应为self.log,这会直接导致运行报错。
  3. 日志逻辑冲突:在test_step中循环log时设置on_epoch=True,会和torchmetrics的内部累积逻辑重复,最终导致指标数值异常。

正确实现方式

1. 初始化指标(保留你的原有逻辑)

在模块的__init__方法中完成每类指标的初始化:

from torchmetrics import F1Score, Precision, Recall
import pytorch_lightning as pl

class YourModel(pl.LightningModule):
    def __init__(self, num_classes):
        super().__init__()
        self.num_classes = num_classes
        # average=None参数指定返回每类的独立指标结果
        self.f1_each = F1Score(task="multiclass", num_classes=num_classes, average=None)
        self.precision_each = Precision(task="multiclass", num_classes=num_classes, average=None)
        self.recall_each = Recall(task="multiclass", num_classes=num_classes, average=None)

2. test_step中累积数据

在test_step里仅负责更新指标的状态,不直接计算或log结果:

def test_step(self, batch, batch_idx):
    x, labels = batch
    preds = self(x)
    
    # 关键:根据模型输出格式调整preds
    # 如果输出是logits,转成类别索引:
    # preds = preds.argmax(dim=1)
    # 或者初始化指标时指定prediction_type="logits"(torchmetrics>=0.10.0支持)
    
    # 将当前batch的数据累积到指标中
    self.precision_each.update(preds, labels)
    self.recall_each.update(preds, labels)
    self.f1_each.update(preds, labels)

3. test_epoch_end中计算并log全局指标

在测试epoch结束时,统一计算整个测试集的每类指标并完成日志记录:

def test_epoch_end(self, outputs):
    # 计算整个测试集的每类指标
    class_precisions = self.precision_each.compute()
    class_recalls = self.recall_each.compute()
    class_f1s = self.f1_each.compute()
    
    # 循环log每类指标
    for i in range(self.num_classes):
        self.log(f'test/precision_{i}', class_precisions[i], prog_bar=True)
        self.log(f'test/recall_{i}', class_recalls[i], prog_bar=True)
        self.log(f'test/f1_{i}', class_f1s[i], prog_bar=True)
    
    # 重置指标状态,避免后续复用出现异常
    self.precision_each.reset()
    self.recall_each.reset()
    self.f1_each.reset()

额外注意事项

  • 确保preds格式匹配指标要求:多分类任务中,默认需要输入类别索引;如果是logits或概率分布,需在初始化指标时指定prediction_type="logits"或prediction_type="probs"。
  • 若使用PyTorch Lightning 2.0+,test_epoch_end可替换为on_test_epoch_end,逻辑完全一致。

内容的提问来源于stack exchange,提问作者Eduard6421

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 21:48:19