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

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 )

问题原因

  1. 重复调用log_hyperparams:main函数中已调用tb_logger.log_hyperparams(cfg),又在on_train_start中再次调用该方法,Lightning仅会处理第一次调用的记录,后续传入的指标映射会被直接忽略。
  2. 指标命名冲突:在log_hyperparams的指标映射中使用了hp/前缀,但Lightning会自动为hparams相关指标添加命名空间,重复前缀导致指标无法和hparams正确关联。
  3. self.hparams未初始化:Lit_train的__init__中未调用self.save_hyperparameters(),导致self.hparams为空,第二次调用log_hyperparams时传入的hparams无效。

修复方案

  1. 移除main函数中的tb_logger.log_hyperparams(cfg),统一在LightningModule内处理超参数和指标映射。
  2. 在Lit_train的__init__中调用self.save_hyperparameters(),确保self.hparams包含需要记录的超参数。
  3. 修改on_train_start中的指标映射,去掉hp/前缀,保持和self.log的指标名称核心一致。
  4. (可选)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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 08:30:06