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
相关产品推荐
相关产品推荐

