PyTorch Lightning Logger异常:测试集仅输出平均指标无中间值
PyTorch Lightning中测试集指标仅记录平均值的问题及解决方法
问题背景
我是PyTorch Lightning新手,正在实现神经网络并绘制各数据集的loss与accuracy曲线。相关代码如下:
def training_step(self, train_batch, batch_idx): X, y = train_batch y_copy = y # Integer y for the accuracy X = X.type(torch.float32) y = y.type(torch.float32) # forward pass y_pred = self.forward(X).squeeze() # accuracy accuracy = Accuracy() acc = accuracy(y_pred, y_copy) # compute loss loss = self.loss_fun(y_pred, y) self.log_dict({'train_loss': loss, 'train_accuracy': acc}, on_step=False, on_epoch=True, prog_bar=True, logger=True) return loss def validation_step(self, validation_batch, batch_idx): X, y = validation_batch X = X.type(torch.float32) # forward pass y_pred = self.forward(X).squeeze() # compute metrics accuracy = Accuracy() acc = accuracy(y_pred, y) loss = self.loss_fun(y_pred, y) self.log_dict({'validation_loss': loss, 'validation_accuracy': acc}, on_step=False, on_epoch=True, prog_bar=True, logger=True) return loss def test_step(self, test_batch, batch_idx): X, y = test_batch X = X.type(torch.float32) # forward pass y_pred = self.forward(X).squeeze() # compute metrics accuracy = Accuracy() acc = accuracy(y_pred, y) loss = self.loss_fun(y_pred, y) self.log_dict({'test_loss': loss, 'test_accuracy': acc}, on_epoch=True, prog_bar=True, logger=True) return loss
训练完成后,执行以下代码读取指标:
metrics = pd.read_csv(f"{trainer.logger.log_dir}/metrics.csv") del metrics["step"] metrics
结果显示,验证集因采用保留法交叉验证(hold out CV)仅输出一组loss和accuracy,这符合预期;但测试集仅输出test_accuracy=0.97这类所有轮次的平均值,无法查看各轮次的中间结果,无法绘制曲线(后续进行K折交叉验证时也需要查看中间结果)。
我发现training_step的日志能正常记录每轮的指标,为何test_step的日志仅输出平均值?如何查看测试集的中间结果?
原因分析
- PyTorch Lightning默认仅在测试阶段结束后计算并记录一次epoch级别的汇总指标,不会像训练/验证阶段那样每轮都记录测试集的指标。因为测试集的核心作用是最终评估模型泛化能力,而非监控训练过程,所以默认行为是输出全局平均值。
- 你的
test_step中使用了on_epoch=True,这会让Lightning在测试完所有batch后计算整个测试集的平均指标并记录一次,而非每个batch或每轮测试的结果。
解决方法
1. 修正测试指标的初始化与日志逻辑
首先,不要在test_step中每次都创建新的Accuracy实例——这会导致指标无法正确累积计算。应该在模型初始化时提前定义:
def __init__(self, ...): super().__init__() # 其他模型层初始化代码 self.test_accuracy = Accuracy()
然后调整test_step的日志参数,如果需要区分不同测试轮次(比如K折交叉验证的不同折),可以开启add_dataloader_idx:
def test_step(self, test_batch, batch_idx): X, y = test_batch X = X.type(torch.float32) y_pred = self.forward(X).squeeze() # 使用提前初始化的指标实例 acc = self.test_accuracy(y_pred, y) loss = self.loss_fun(y_pred, y) self.log_dict( {'test_loss': loss, 'test_accuracy': acc}, on_step=False, on_epoch=True, prog_bar=True, logger=True, add_dataloader_idx=True # 多折测试时,自动为指标添加折数标识 ) return loss
2. 手动遍历每轮模型测试并记录
如果需要查看训练过程中每轮模型在测试集上的表现,可以遍历训练时保存的checkpoint文件,逐个加载并测试:
import glob import pandas as pd from pytorch_lightning import Trainer # 匹配所有保存的epoch checkpoint checkpoint_paths = sorted(glob.glob('lightning_logs/version_0/checkpoints/epoch=*.ckpt')) test_metrics = [] for ckpt_path in checkpoint_paths: # 从文件名解析epoch数 epoch_num = int(ckpt_path.split('epoch=')[-1].split('.ckpt')[0]) # 加载对应epoch的模型 model = YourLightningModule.load_from_checkpoint(ckpt_path) trainer = Trainer(logger=False) # 运行测试 results = trainer.test(model, datamodule=your_test_datamodule)[0] # 绑定epoch信息 results['epoch'] = epoch_num test_metrics.append(results) # 转换为DataFrame方便绘图 test_metrics_df = pd.DataFrame(test_metrics) print(test_metrics_df)
3. 用回调实现训练中定期测试
如果想在训练过程中每N轮自动运行测试并记录指标,可以自定义回调:
from pytorch_lightning.callbacks import Callback class TestEveryNEpoch(Callback): def __init__(self, test_datamodule, every_n_epochs=1): self.test_datamodule = test_datamodule self.every_n_epochs = every_n_epochs def on_epoch_end(self, trainer, pl_module): if trainer.current_epoch % self.every_n_epochs == 0: # 静默运行测试 results = trainer.test(pl_module, datamodule=self.test_datamodule, verbose=False)[0] # 为指标添加epoch标识并记录 for metric_name, value in results.items(): trainer.logger.log_metrics( {f'test_{metric_name}_epoch_{trainer.current_epoch}': value}, step=trainer.current_epoch ) # 训练时添加该回调 trainer = Trainer(callbacks=[TestEveryNEpoch(your_test_datamodule)]) trainer.fit(model, datamodule=your_train_datamodule)
内容的提问来源于stack exchange,提问作者Dylan Nico
相关产品推荐
相关产品推荐

