我正在笔记本上工作。我用下面的代码做了一些初步的实验。
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() 函数。
我会尽快发布我的完整笔记本。你也可以在那里检查。