You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 03:15:19