PyTorch Lightning自定义hp/metrics无法正常工作求助
Lightning多指标记录到TensorBoard hparams失效问题排查
问题背景
Lightning.pytorch支持通过hp_metric相关配置将自定义指标(如损失函数)记录到TensorBoard的hparams板块,方便用户通过筛选框按阈值筛选实验。但配置多指标后功能完全失效,TensorBoard的hparams板块显示异常,相关代码如下:
import torch import lightning.pytorch as pl from train import do_training, Lit_train from lightning.pytorch.loggers import TensorBoardLogger def main(cfg: DictConfig): ....... tb_logger = TensorBoardLogger("tb_logs", name="Experiment", default_hp_metric=False) tb_logger.log_hyperparams(cfg) model = Lit_train(model, .....,) trainer = pl.Trainer( .....,logger=tb_logger) trainer.fit( model, train_loader, valid_loader) if __name__ =="__main__": main() class Lit_train(pl.LightningModule): def __init__(self, model,.....): ..... def on_train_start(self): self.logger.log_hyperparams(self.hparams, {"hp/metric_1": 0, "hp/metric_2": 0}) def training_step(self, batch, batch_idx): c, targets = batch rollout = self.model(c) loss = loss_fun(rollout, targets) return loss def validation_step(self, batch, batch_idx): c, targets = batch rollout = self.model(c) loss = loss_fun(rollout, targets) self.log("hp/metric_1",loss ) self.log("hp/metric_2",loss )
问题原因
- 重复调用
log_hyperparams:main函数中已调用tb_logger.log_hyperparams(cfg),又在on_train_start中再次调用该方法,Lightning仅会处理第一次调用的记录,后续传入的指标映射会被直接忽略。 - 指标命名冲突:在
log_hyperparams的指标映射中使用了hp/前缀,但Lightning会自动为hparams相关指标添加命名空间,重复前缀导致指标无法和hparams正确关联。 self.hparams未初始化:Lit_train的__init__中未调用self.save_hyperparameters(),导致self.hparams为空,第二次调用log_hyperparams时传入的hparams无效。
修复方案
- 移除
main函数中的tb_logger.log_hyperparams(cfg),统一在LightningModule内处理超参数和指标映射。 - 在
Lit_train的__init__中调用self.save_hyperparameters(),确保self.hparams包含需要记录的超参数。 - 修改
on_train_start中的指标映射,去掉hp/前缀,保持和self.log的指标名称核心一致。 - (可选)
self.log时指定prog_bar=False,避免进度条重复显示无关指标。
修复后的代码示例
import torch import lightning.pytorch as pl from train import do_training, Lit_train from lightning.pytorch.loggers import TensorBoardLogger from omegaconf import DictConfig def main(cfg: DictConfig): # 初始化TensorBoardLogger,关闭默认hp_metric tb_logger = TensorBoardLogger("tb_logs", name="Experiment", default_hp_metric=False) # 初始化模型,传入必要参数 model = Lit_train(model=your_model_instance, cfg=cfg, ...) # 初始化Trainer trainer = pl.Trainer(..., logger=tb_logger) trainer.fit(model, train_loader, valid_loader) if __name__ == "__main__": main() class Lit_train(pl.LightningModule): def __init__(self, model, cfg, ...): super().__init__() self.model = model # 保存超参数到self.hparams,ignore参数可排除不需要记录的对象(如模型实例) self.save_hyperparameters(cfg, ignore=['model']) def on_train_start(self): # 仅调用一次log_hyperparams,传入超参数和指标初始值(无hp/前缀) self.logger.log_hyperparams( self.hparams, {"metric_1": 0.0, "metric_2": 0.0} ) def training_step(self, batch, batch_idx): c, targets = batch rollout = self.model(c) loss = loss_fun(rollout, targets) return loss def validation_step(self, batch, batch_idx): c, targets = batch rollout = self.model(c) loss = loss_fun(rollout, targets) # 记录指标时添加hp/前缀,确保和hparams映射关联 self.log("hp/metric_1", loss, logger=True, prog_bar=False) self.log("hp/metric_2", loss, logger=True, prog_bar=False)
内容的提问来源于stack exchange,提问作者Andrey Vlasenko
相关产品推荐
相关产品推荐

