PyTorch Lightning:如何通过WandbLogger记录W&B直方图?
使用WandbLogger记录直方图的方法
方法一:通过LightningModule的log方法自动记录
在你的LightningModule的训练、验证或测试步骤中,直接使用self.log()传入张量数据即可,WandbLogger会自动将其转换为直方图展示:
import torch import pytorch_lightning as pl from pytorch_lightning.loggers import WandbLogger class MyModel(pl.LightningModule): def __init__(self): super().__init__() self.layer = torch.nn.Linear(10, 2) self.loss_fn = torch.nn.CrossEntropyLoss() def training_step(self, batch, batch_idx): x, y = batch pred = self(x) loss = self.loss_fn(pred, y) # 记录模型层权重的直方图 self.log("layer_weights/hist", self.layer.weight, prog_bar=False) # 记录预测结果的直方图 self.log("train_predictions/hist", pred, prog_bar=False) return loss
方法二:调用W&B原生API自定义记录
如果需要更精细的控制(比如自定义直方图的bins、分组逻辑),可以通过WandbLogger获取W&B的运行实例,使用wandb.log()配合wandb.Histogram来记录:
class MyModel(pl.LightningModule): def validation_step(self, batch, batch_idx): x, y = batch pred = self(x) # 获取W&B运行实例 wandb_run = self.logger.experiment # 用wandb.Histogram包装数据,自定义记录 wandb_run.log({ "val_pred_hist": wandb.Histogram(pred.detach().cpu().numpy()), "val_label_hist": wandb.Histogram(y.detach().cpu().numpy()) })
批量记录所有可训练参数的直方图
如果需要一次性记录模型所有可训练参数的直方图,可以在 epoch 结束时遍历参数:
def on_train_epoch_end(self): wandb_run = self.logger.experiment for name, param in self.named_parameters(): if param.requires_grad: wandb_run.log({ f"params/{name}_hist": wandb.Histogram(param.detach().cpu().numpy()) })
内容的提问来源于stack exchange,提问作者Rylan Schaeffer
相关产品推荐
相关产品推荐

