如何在PyTorch Lightning的test_step中添加每类精确率、召回率、F1指标?
PyTorch Lightning test_step中每类指标计算的问题与解决方法
你代码中的错误点
- 指标更新逻辑错误:直接调用
self.precision_each(preds, labels)只会计算当前batch的指标,无法累积整个测试集的结果。torchmetrics的指标类需要通过update方法逐步累积数据,最终用compute得到全局统计结果。 - 拼写错误:代码中
self.logger.expe是笔误,应为self.log,这会直接导致运行报错。 - 日志逻辑冲突:在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
相关产品推荐
相关产品推荐

