尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch Lightning TPU 性能剖析进阶:用 XLAProfiler 在 TensorBoard 中定位 Cloud TPU 训练瓶颈

PyTorch Lightning TPU 性能剖析进阶:用 XLAProfiler 在 TensorBoard 中定位 Cloud TPU 训练瓶颈 PyTorch Lightning TPU 性能剖析进阶用 XLAProfiler 在 TensorBoard 中定位 Cloud TPU 训练瓶颈【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning导读本文面向希望在 Cloud TPU 上训练并优化性能的 PyTorch Lightning 用户系统讲解如何通过内置的XLAProfiler剖析 TPU 模型、定位训练瓶颈并把剖析日志捕获到 TensorBoard 中进行可视化分析。读完本文你将掌握XLAProfiler的接入方式、其底层运行原理对应 xla.py 源码实现以及从启动 TensorBoard 到完成 profile 捕获的完整实操流程能够独立对 TPU 训练任务开展性能诊断。一、为什么需要 TPU 级剖析Trainer(profilersimple)只能统计 Lightning 内部各 hook 的耗时Trainer(profileradvanced)基于 PythoncProfile剖析函数级调用但当你运行在 Cloud TPUXLA 设备上时计算图由 XLA 编译执行Python 侧的计时无法反映真实设备执行情况。此时需要的是 docs/source-pytorch/tuning/profiler_advanced.rst 所讲解的方案使用torch_xla.debug.profiler提供的 XLA 服务器能力抓取设备端 trace再通过 TensorBoard 的 Profile 插件查看硬件利用率、kernel 耗时与数据搬运等关键信息。本文属于 profiler 系列中的Advanced章节面向「需要剖析 TPU 模型以发现并改进性能瓶颈」的用户更基础的训练循环剖析见 profiler_basic.rst自定义 profiler 的专家级内容见 profiler_expert.rst。二、前置条件确保 XLA 环境可用XLAProfiler依赖torch_xla。从 xla.py 的源码可以看到构造XLAProfiler时若环境不可用会立即抛出异常if not _XLA_AVAILABLE: raise ModuleNotFoundError(str(_XLA_AVAILABLE))其中_XLA_AVAILABLE来自lightning.fabric.accelerators.xla。因此在使用本方案前请确保训练环境为 Cloud TPU VM或具备torch_xla的 TPU 环境按官方 TPU 指引完成所需的安装与初始化如pip install cloud-tpu-client、libtpu相关运行时依赖等此处不再展开具体可对照官方 Cloud TPU 的 PyTorch/XLA 性能剖析安装说明。三、在 Lightning 中启用 XLAProfiler3.1 直接实例化与文档示例一致最简单的接入方式如下from lightning.pytorch import Trainer from lightning.pytorch.profilers import XLAProfiler profiler XLAProfiler(port9001) trainer Trainer(profilerprofiler)3.2 使用字符串快捷方式除了直接传入实例Trainer的profiler参数还支持字符串形式。在 setup.py 的_init_profiler中内置了四类 profiler 的字符串映射PROFILERS { simple: SimpleProfiler, advanced: AdvancedProfiler, pytorch: PyTorchProfiler, xla: XLAProfiler, }因此可以直接写trainer Trainer(profilerxla, acceleratortpu, devicesauto)这一用法同样被测试用例覆盖——见 test_xla_profiler.py其中断言了trainer.profiler确实是XLAProfiler实例。3.3 关于 port 参数的说明文档示例中显式传入port9001这与后面 TensorBoard 的 9001 端口保持一致便于通过浏览器http://localhost:9001/#profile统一访问。需要留意一个细节从 xla.py 的源码签名看XLAProfiler.__init__的默认端口是port: int 9012并非文档中提到的 9001。因此如果你不传portXLA Profiler 服务器默认监听9012若希望浏览器里直接填localhost:9001捕获请如文档示例那样显式传入port9001并保证该端口未被占用。源码的 docstring 也明确说明端口非法或被占用时会抛出异常An exception is raised if the provided port is invalid or busy。四、XLAProfiler 的底层工作原理理解XLAProfiler做了什么有助于你正确使用它。对照 xla.py 源码4.1 被记录的动作集合STEP_FUNCTIONS {validation_step, test_step, predict_step} RECORD_FUNCTIONS { training_step, backward, validation_step, test_step, predict_step, }RECORD_FUNCTIONS定义了会被记录的核心动作STEP_FUNCTIONS定义了哪些动作按「步」记录带 step 序号其余动作按普通Trace记录。4.2 start / stop 的调用链Profiler基类profiler.py通过profile()上下文管理器串联start()与stop()XLAProfiler重写了这两个方法首次触发记录动作时xp.start_server(self.port)启动 XLA Profiler 服务器之后所有记录挂到该服务器上等待外部捕获对training_step、backward等动作进入xp.Trace(action_name)对validation_step、test_step、predict_step则使用xp.StepTrace(action_name, step_numstep)并通过_get_step_num()维护每个动作的步号计数stop()时调用对应 trace 对象的__exit__结束记录。由此可见Lightning 训练循环中的这些核心动作会被自动打点你只需在外部发起捕获请求即可拿到设备侧的执行 trace。4.3 Trainer 如何装配 profiler在 trainer.py 的__setup_profiler中Trainer会在训练开始时调用self.profiler._lightning_module proxy(self.lightning_module) self.profiler.setup(stageself.state.fn, local_ranklocal_rank, log_dirself.log_dir)即 profiler 会随训练阶段自动完成 setup你无需手动启动服务器只要保证捕获期间训练代码正在运行即可。五、把剖析日志捕获到 TensorBoard完整步骤下面按文档的流程完整复现从安装到查看剖析结果的全部操作。第 0 步完成 Cloud TPU 所需的安装请参考 Cloud TPU 官方关于「PyTorch/XLA 在 TPU VM 上的性能剖析」的安装指引确保torch_xla及 profiling 相关依赖包括 TensorBoard 插件都已就绪。这一步属于环境准备缺失会导致后续XLAProfiler构造或服务器启动失败。第 1 步启动 TensorBoard 服务器在终端中执行tensorboard --logdir ./tensorboard --port 9001然后在浏览器中打开http://localhost:9001/#profile打开后即可看到 TensorBoard 的 Profile 页面。注意--port 9001应与代码中XLAProfiler(port9001)保持一致避免端口错配。第 2 步捕获 profile当你要剖析的代码已经在运行中时在 Profile 页面执行点击CAPTURE PROFILE按钮在 Profile Service URL 一栏填入localhost:9001即 XLA Profiler 服务器的地址输入期望的剖析时长单位毫秒点击CAPTURE开始捕获。捕获过程由 XLA Profiler 服务器收集设备端 trace 并写入 TensorBoard 日志目录随后页面会展示剖析结果。第 3 步保持代码持续运行捕获期间不要停止正在运行的训练代码。文档特别强调剖析时长最好大于一个 step 的耗时这样捕获到的 trace 能覆盖完整的训练迭代获得的性能洞察更有代表性。若剖析时长过短可能只截到某个 step 的片段难以定位瓶颈。第 4 步查看剖析日志捕获完成后页面会自动刷新你可以通过左上角的Tools下拉菜单浏览各项剖析结果例如 kernel 耗时统计、trace view、设备利用率等从而定位是计算 kernel、数据加载还是通信导致的瓶颈。六、进阶程序化捕获 trace来自测试用例的做法除了手动在 TensorBoard 页面点击捕获仓库的测试用例 test_xla_profiler.py 展示了程序化捕获方式——通过torch_xla的xp.traceAPI 直接抓取import torch_xla.debug.profiler as xp # 训练在独立进程中持续运行profilerxla logdir str(tmp_path) xp.trace( flocalhost:{port}, logdir, duration_ms2000, # 剖析时长 2000ms num_tracing_attempts5, # 最多尝试 5 次 delay_ms1000, # 延迟 1000ms 后再开始 )该测试随后断言日志目录下生成了plugins/profile/*/*.xplane.pb这样的 trace 文件这些文件正是 TensorBoard Profile 插件读取的数据源。对于需要自动化性能回归测试或批量收集 trace 的场景可以借鉴这一写法。七、注意事项与常见坑不要用torch.profiler.profile手动包裹Trainer.fit()在 profiler_basic.rst 中有明确警告——手动包裹 Trainer 方法会因 PyTorch Profiler 的上下文管理与 Lightning 内部训练循环不兼容而导致意外崩溃或晦涩报错。应始终通过Trainer(profiler...)传入或在需要定制时使用PyTorchProfiler类。端口一致性XLAProfiler(port...)、tensorboard --port ...、Profile Service URL 三者需指向同一端口默认端口为 9012源码默认值文档示例统一使用 9001。捕获时机服务器在第一次触发被记录动作时才启动见 xla.py 的if not self._start_trace分支因此务必在训练运行中发起捕获。仅适用于 TPU/XLA 环境非 TPU 环境构造XLAProfiler会直接抛出ModuleNotFoundError。八、延伸阅读剖析入门训练循环瓶颈profiler_basic.rst剖析总览索引profiler.rst自定义 profiler 与剖析指定代码段profiler_expert.rst核心实现xla.py、profiler.py、profiler 入口汇总见init.pyprofiler 字符串映射与装配逻辑setup.py、trainer.pyTPU 剖析测试用例test_xla_profiler.py【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表