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

PyTorch Lightning训练中如何追踪最优指标分数

在PyTorch Lightning中追踪Epoch级别的损失与性能指标

我正在学习PyTorch Lightning,想要追踪每个epoch的损失和性能指标分数,用来保存最优模型、绘制图表或者获取最高性能,但目前只能在training_step方法中获取批次级别的损失和性能分数,代码示例如下:

class Model(pl.LightningModule):
    ...

    def training_step(self, batch, batch_idx):
        ...

        probs = self.forward(input_ids, labels, output_ids)
        loss = self.loss(probs, labels)
        
        # 计算BLEU分数
        predictions = self.decode_prediction(probs, input_tokens)
        output_sequences = [[x] for x in output_sequences]
        blue_score = self.blue_score(predictions, output_sequences)
        
        self.log_dict(
            {'train_loss': loss, 'blue_score': blue_score},
            on_step=True,
            on_epoch=True,
            prog_bar=True
        )

        return loss

经过调研后找到了解决方案:因为设置了on_epoch = True,默认Logger会自动累加每个epoch的损失和指标分数。可以编写一个继承自Torch Lightning Logger类的自定义Logger来收集这些epoch级别的数据,代码示例如下:

import collections
from pytorch_lightning.loggers import Logger
from pytorch_lightning.utilities.rank_zero import rank_zero_only

class HistoryLogger(Logger):
    def __init__(self):
        super().__init__()
        self.history = collections.defaultdict(list)

    @property
    def name(self):
        return "HistoryLogger"

    @property
    def version(self):
        return "1.0"

    @rank_zero_only
    def log_metrics(self, metrics, step):            
        for metric_name, metric_value in metrics.items():
            self.history[metric_name].append(metric_value)
        return

logger = HistoryLogger()

trainer = Trainer(logger=logger)
...

# 访问历史数据
logger.history

内容的提问来源于stack exchange,提问作者jupyter

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 08:50:08