如何将PyTorch Lightning Profiler与TensorBoard集成?
问题原因
- 你将
torch.profiler.profile上下文管理器放在了training_step内部,每次迭代step都会创建一个全新的profiler实例,schedule设定的wait/warmup/active全局计数逻辑完全无法正常执行,也不会触发trace写入逻辑,因此无法生成对应profiler日志。
解决方案
优先使用PyTorch Lightning内置封装的PyTorchProfiler,它已经原生适配了训练生命周期,支持直接导出TensorBoard可读的profiler日志,不需要自己手动管理profiler生命周期,配置方式和原生PyTorch profiler完全兼容:
from pytorch_lightning.profilers import PyTorchProfiler import torch # 初始化profiler,参数和原生torch.profiler完全对齐 profiler = PyTorchProfiler( activities=[torch.profiler.ProfilerActivity.CPU], schedule=torch.profiler.schedule(wait=1, warmup=1, active=2, repeat=1), with_stack=True, on_trace_ready=torch.profiler.tensorboard_trace_handler('./logs'), record_shapes=True ) # 传入Trainer即可,不需要修改原有training_step代码 trainer = Trainer( profiler=profiler, # 其余你原有Trainer的配置保持不变 ) # 正常执行训练 trainer.fit(你的模型实例, 训练数据集加载器)
配置完成后Lightning会自动管理profiler的启停和step调用,生成的profiler日志会直接写入你指定的./logs目录,直接用TensorBoard加载该目录即可查看profiler相关面板。
如果你坚持要手动使用原生torch profiler实现,则需要将profiler的初始化和销毁放在训练全局钩子中,不要在training_step中重复创建实例:
class 你的模型类(LightningModule): def on_train_start(self): # 训练启动时全局初始化一次profiler self.profiler = torch.profiler.profile( activities=[torch.profiler.ProfilerActivity.CPU], schedule=torch.profiler.schedule( wait=1, warmup=1, active=2, repeat=1), with_stack=True, on_trace_ready=torch.profiler.tensorboard_trace_handler('./logs'), ) self.profiler.start() def training_step(self, train_batch, batch_idx): x, y = train_batch x = x.float() logits = self.forward(x) loss = self.loss_fn(logits, y) # 每步仅调用step方法 self.profiler.step() return loss def on_train_end(self): # 训练结束时停止profiler,触发最后的日志写入 self.profiler.stop()
内容的提问来源于stack exchange,提问作者Madara
相关产品推荐
相关产品推荐

