【问题标题】:How to get total test accuracy for pytorch lightning?如何获得 pytorch 闪电的总测试精度?
【发布时间】:2023-01-23 13:00:13
【问题描述】:

trainer.test 方法如何用于获得所有批次的总准确度?

我知道我可以实施model.test_step,但这仅适用于单个批次。我需要整个数据集的准确性。我可以使用torchmetrics.Accuracy 来累积准确率。但是,将它们结合起来并获得总准确度的正确方法是什么?由于分批测试分数不是很有用,model.test_step 无论如何应该返回什么?我可以以某种方式破解它,但令我惊讶的是我在互联网上找不到任何示例来演示如何使用 pytorch-lightning 本机方式获得准确性。

【问题讨论】:

    标签: pytorch-lightning


    【解决方案1】:

    您可以在这里(https://pytorch-lightning.readthedocs.io/en/stable/extensions/logging.html#automatic-logging)看到log 中的on_epoch 参数在纪元结束时自动累积并记录。这样做的正确方法是:

    from torchmetrics import Accuracy
    
    def validation_step(self, batch, batch_idx): 
        x, y = batch 
        preds = self.forward(x) 
        loss = self.criterion(preds, y) 
        accuracy = Accuracy()
        acc = accuracy(preds, y)
        self.log('accuracy', acc, on_epoch=True)
        return loss 
    

    如果你想要一个自定义的缩减函数,你可以使用reduce_fx参数来设置它,默认是torch.mean()。 log()可以从你LightningModule中的任何方法调用

    【讨论】:

    • 当它不知道 batchsize 时,它​​如何累加? (批次可能不相等,或者至少最后一个批次的大小不同)。什么是平均法?我的意思也是测试,即test_step。它还会起作用吗?
    • 根据您在上面评论中的问题更新了答案
    • 谢谢。我做了一个测试,它似乎工作。实际上,缩减方法不能是普通的mean,因为你不能只是平均批处理精度。但我想它宁愿使用完整的 Accuracy 对象,并且该对象知道它应该如何减少。
    • 要了解您的平均值,您应该查看 log 和 Accuracy() (torchmetrics.readthedocs.io/en/latest/classification/…)。它们都可以以不同的方式进行平均。
    【解决方案2】:

    我正在笔记本上工作。我用下面的代码做了一些初步的实验。

    def test_step(self, batch, batch_idx):
        x, y = batch
        logits = self(x)
        self.test_acc(logits, y)
        self.log('test_acc', self.test_acc, on_step=False, on_epoch=True)
    

    调用后打印出格式良好的文本

    model = Cifar100Model()
    trainer = pl.Trainer(max_epochs=1, accelerator='cpu')
    trainer.test(model, test_dataloader)
    

    此打印的 test_acc 0.008200000040233135

    我尝试验证打印值是否实际上是测试数据批次的平均值。通过修改 test_step 如下:

    def test_step(self, batch, batch_idx):
        x, y = batch
        logits = self(x)
        self.test_acc(logits, y)
        self.log('test_acc', self.test_acc, on_step=False, on_epoch=True)
    
        preds = logits.argmax(dim=-1)
        acc = (y == preds).float().mean()
        print(acc)
    

    然后再次运行 trainer.test() 。这次打印出以下值:
    张量(0.0049)
    张量(0.0078)
    张量(0.0088)
    张量(0.0078)
    张量(0.0122)
    平均它们得到我:0.0083 这非常接近 test_step() 打印的值。

    这个解决方案背后的逻辑是我在

    self.log('test_acc', self.test_acc, on_step=False, on_epoch=True)
    

    on_epoch = True,我使用了 TorchMetric 类,平均值由 PL 计算,自动使用 metric.compute() 函数。

    我会尽快发布我的完整笔记本。你也可以在那里检查。

    【讨论】:

      猜你喜欢
      • 2021-07-04
      • 2019-02-10
      • 2021-07-05
      • 2015-01-27
      • 2022-06-27
      • 2021-01-13
      • 2012-01-13
      • 2015-02-12
      • 2021-04-24
      相关资源
      最近更新 更多