如何在PyTorch Lightning中每N步打印训练及验证集的损失与准确率
每N个训练步骤打印训练与验证指标的简洁实现
下面是一个轻量、易理解的PyTorch Lightning实现方案,满足你每经过N个训练批次就打印训练/验证loss、准确率的需求:
核心实现代码
import pytorch_lightning as pl from torchmetrics import Accuracy import torch from torch.utils.data import DataLoader, TensorDataset class MyModel(pl.LightningModule): def __init__(self, log_every_n_steps=10): super().__init__() # 示例模型层(替换成你的实际模型) self.fc = torch.nn.Linear(10, 2) # 训练指标 self.train_acc = Accuracy(task="multiclass", num_classes=2) self.train_loss_buffer = [] # 验证指标 self.val_acc = Accuracy(task="multiclass", num_classes=2) self.val_loss_buffer = [] # 配置触发间隔 self.log_interval = log_every_n_steps self.step_counter = 0 def forward(self, x): return self.fc(x) def training_step(self, batch, batch_idx): x, y = batch logits = self(x) loss = torch.nn.functional.cross_entropy(logits, y) # 更新训练指标缓存 self.train_acc(logits, y) self.train_loss_buffer.append(loss.item()) # 计数器递增,达到间隔触发日志与验证 self.step_counter += 1 if self.step_counter % self.log_interval == 0: self._run_log_and_validation() return loss def _run_log_and_validation(self): # 计算训练批次的平均指标 avg_train_loss = sum(self.train_loss_buffer) / len(self.train_loss_buffer) avg_train_acc = self.train_acc.compute() # 切换模型到评估模式,避免影响训练 self.eval() self.val_acc.reset() self.val_loss_buffer.clear() # 手动遍历验证集计算指标(无需依赖trainer的复杂流程) with torch.no_grad(): for val_batch in self.val_dataloader(): x_val, y_val = val_batch val_logits = self(x_val) val_loss = torch.nn.functional.cross_entropy(val_logits, y_val) self.val_acc(val_logits, y_val) self.val_loss_buffer.append(val_loss.item()) # 计算验证集平均指标 avg_val_loss = sum(self.val_loss_buffer) / len(self.val_loss_buffer) avg_val_acc = self.val_acc.compute() # 切回训练模式 self.train() # 打印到终端 print(f"\n=== Step {self.step_counter} Metrics ===") print(f"Train Loss: {avg_train_loss:.4f} | Train Acc: {avg_train_acc:.4f}") print(f"Val Loss: {avg_val_loss:.4f} | Val Acc: {avg_val_acc:.4f}\n") # 记录到日志系统(支持TensorBoard、WandB等) self.log("train/loss", avg_train_loss, step=self.step_counter) self.log("train/acc", avg_train_acc, step=self.step_counter) self.log("val/loss", avg_val_loss, step=self.step_counter) self.log("val/acc", avg_val_acc, step=self.step_counter) # 重置训练指标缓存,准备下一轮计数 self.train_acc.reset() self.train_loss_buffer.clear() def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr=1e-3) # 替换成你的实际验证数据集加载逻辑 def val_dataloader(self): x_val = torch.randn(100, 10) y_val = torch.randint(0, 2, (100,)) return DataLoader(TensorDataset(x_val, y_val), batch_size=10)
关键说明
- 计数器控制触发时机:用
step_counter跟踪训练批次数量,每达到设定的log_interval就执行一次验证与日志打印 - 轻量验证逻辑:手动遍历验证集计算指标,避免调用
trainer.validate带来的额外配置复杂度 - 模式切换与梯度隔离:切换模型到
eval模式并使用torch.no_grad(),防止验证过程影响训练状态和产生不必要的梯度计算 - 双端日志输出:同时在终端打印直观信息,以及记录到Lightning的日志系统,支持后续可视化分析
- 指标重置:每次触发后清空训练指标缓存,确保下一轮计数的指标是全新的批次数据
内容的提问来源于stack exchange,提问作者Michael D
相关产品推荐
相关产品推荐

