首页 / 资讯中心 / 文章详情

使用 torch.compile 编译优化器(Adam)加速 PyTorch 训练:实战指南

使用 torch.compile 编译优化器(Adam)加速 PyTorch 训练:实战指南 ★ FEATURED ARTICLE
示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载优化器负责更新模型的每一个参数在大模型训练中往往成为性能瓶颈。本文基于 PyTorch 官方教程仓库中的 compiling_optimizer.rst手把手演示如何用torch.compile编译优化器的step()方法并通过torch.utils.benchmark量化 GPU 上的性能提升。读完本文你将掌握编译优化器的完整实操流程、基准测试的正确姿势以及优化器与 LR Scheduler 搭配时的重编译陷阱。为什么优化器会成为训练瓶颈在任何一个深度学习模型的训练循环中优化器承担着最关键的工作读取每个参数的梯度按照学习率等超参数更新参数值。对于大模型参数动辄数千万乃至数十亿这意味着逐参数更新带来大量内存读写每一步更新都要读取参数与梯度、写入更新后的参数属于典型的内存密集型memory-bound操作每个参数的更新操作独立成 kernel在 eager 模式下PyTorch 为每个参数更新操作单独启动 kernelkernel launch 开销与 Python 解释器开销叠加显著放大耗时优化器状态如 Adam 的一阶/二阶矩进一步加重负担Adam 类优化器还需维护额外的状态张量读写量成倍增加。正因如此当模型变大后optimizer.step()在训练性能中的占比会越来越高成为值得单独优化的对象。本教程的核心思路就是把step()包进torch.compile让底层编译器把一系列逐参数更新操作融合成更少的 kernel从而减少内存往返与启动开销。注意本教程需要PyTorch 2.2.0 或更高版本torch.compile于 PyTorch 2.0 引入编译优化器的完整支持与配套基准代码在 2.2 中可用。同时torch.compile仅支持compute capability 7.0Volta 及以上的 CUDA 设备。模型搭建只关心参数数量教程选择了一个由 10 层Linear组成的简单顺序模型。作者特别强调由于我们只基准测试优化器本身模型的具体结构无关紧要——优化器的性能只取决于参数数量。import torch model torch.nn.Sequential( *[torch.nn.Linear(1024, 1024, False, devicecuda) for _ in range(10)] ) input torch.rand(1024, devicecuda) output model(input) output.sum().backward()关键点说明torch.nn.Linear(1024, 1024, False, devicecuda)中第三个参数biasFalse去掉了偏置项让每一层只含权重矩阵参数结构更干净先执行一次前向传播再通过output.sum().backward()完成反向传播为优化器填充梯度。这一步是必须的——没有梯度opt.step()无事可做10 层 × 1024×1024 权重共约 1000 万参数规模适中足以让优化器更新开销成为可观测的测量对象。设置并运行优化器基准设备能力检查由于torch.compile对设备有硬性要求教程在入口处就做了防护在不受支持的设备上干净地退出# exit cleanly if we are on a device that doesnt support torch.compile if torch.cuda.get_device_capability() (7, 0): print(Exiting because torch.compile is not supported on this device.) import sys sys.exit(0)torch.cuda.get_device_capability()返回当前 CUDA 设备的 (major, minor) 计算能力元组例如(8, 0)对应 Ampere 架构。低于(7, 0)Volta 之前的设备直接退出。编译优化器的 step()接着创建 Adam 优化器并定义一个用torch.compile装饰的包装函数把step()包进去opt torch.optim.Adam(model.parameters(), lr0.01) torch.compile(fullgraphFalse) def fn(): opt.step()这里两个细节值得展开fullgraphFalse这是torch.compile的默认设置表示允许图断裂graph break。TorchDynamo 在追踪时若遇到难以捕获的 Python 代码会中断编译、退回 eager 执行这部分代码然后继续编译。fullgraphTrue则会在遇到第一个图断裂时直接报错。对本例而言opt.step()内部逻辑可以被完整捕获fullgraphFalse只是保持默认的容错行为。关于图断裂的深入讨论可参考 torch_compile_tutorial.py捕获的边界是 Python 函数torch.compile是装饰器作用于任意 Python 函数。编译发生时 TorchDynamo 对fn的字节码进行追踪捕获其中的 PyTorch 算子序列交给 TorchInductor 生成融合后的底层 kernelCUDA 下通常是 Triton kernel后续调用直接复用编译产物。基准测试辅助函数与测量教程使用torch.utils.benchmark提供的Timer与blocked_autorange来获得稳定、可统计的耗时# Lets define a helpful benchmarking function: import torch.utils.benchmark as benchmark def benchmark_torch_function_in_microseconds(f, *args, **kwargs): t0 benchmark.Timer( stmtf(*args, **kwargs), globals{args: args, kwargs: kwargs, f: f} ) return t0.blocked_autorange().mean * 1e6关于该测量方式的原理仓库中的 benchmark.py 给出了详细说明与标准库timeit不同torch.utils.benchmark.Timer会自动处理CUDA 同步eager 的 kernel 是异步发射的不同步只能量到发射时间而非真实执行时间blocked_autorange()会先通过递增的单次运行次数找到合适规模这一过程本身起到warmup作用再连续多次测量直至累计时长达到目标默认至少 0.2 秒可用min_run_time调整返回的Measurement对象带有mean、median等统计量便于评估测量可靠性这里取mean * 1e6把秒换算成微秒us。完整测量流程# Warmup runs to compile the function for _ in range(5): fn() eager_runtime benchmark_torch_function_in_microseconds(opt.step) compiled_runtime benchmark_torch_function_in_microseconds(fn) assert eager_runtime compiled_runtime print(feager runtime: {eager_runtime}us) print(fcompiled runtime: {compiled_runtime}us)几个容易忽略但至关重要的点必须先 warmup 再测量torch.compile的第一次调用会触发完整编译流程耗时远高于后续调用详见 torch_compile_tutorial.py 中首次编译耗时偏大的演示。教程用 5 次循环预热让编译产物缓存就位后再计时分别测量 eager 与 compiledeager 基线直接测opt.step编译版本测包装函数fnassert eager_runtime compiled_runtime是门禁它确保在编译确实带来加速时才继续若某个环境如编译产物异常、测量受干扰下编译版本反而更慢程序会在此处显式失败避免输出误导性结论结果具有机器相关性正如文档明确提示的Depending on what machine you are using, your exact results may vary加速比取决于 GPU 型号、驱动、编译缓存等因素示例数值仅作量级参考。示例结果教程给出的单次参考输出为Eager runtime约 747.24 usCompiled runtime约 392.07 us约 1.9 倍的提升。提速来源主要是TorchInductor 将 Adam 更新中原本逐参数串行执行的多个 pointwise 算子计算梯度一阶矩、二阶矩、偏差修正、参数更新等融合成更少的 kernel从而大幅减少内存往返与 kernel 启动开销——这与 tuning_guide.py 中算子融合一节描述的原理一致pointwise 算子通常受内存带宽限制每融合一个算子就少一次完整的数据加载与回写。进阶编译优化器与 LR Scheduler 搭配基础教程之外仓库中的配套脚本 compiling_optimizer_lr_scheduler.py 展示了真实训练中更常见的场景把编译后的优化器与学习率调度器一起使用该示例要求PyTorch 2.3.0 或更高版本。# !!! IMPORTANT !!! Wrap the lr in a Tensor if we are pairing the # the optimizer with an LR Scheduler. # Without this, torch.compile will recompile as the value of the LR # changes. opt torch.optim.Adam(model.parameters(), lrtorch.tensor(0.01)) sched torch.optim.lr_scheduler.LinearLR(opt, total_iters5) torch.compile(fullgraphFalse) def fn(): opt.step() sched.step() # Warmup runs to compile the function for _ in range(5): fn() print(opt.param_groups[0][lr])这里有一个决定成败的细节把学习率包装成torch.Tensorlrtorch.tensor(0.01)。原因是torch.compile会在每次调用时用 guard 校验输入状态是否与已编译版本一致。如果lr是普通 Python 浮点数调度器每次step()都会修改其值触发guard 失败导致函数在每一次迭代都重新编译——编译时间反而成为新的瓶颈。将lr包成 Tensor 后其值变化以张量数据的形式参与计算不再触发 guard 层面的重编译。该脚本还专门演示了如何用日志验证这一点在非 Tensor 场景下开启重编译日志就能观察到调度器步进导致的重复编译# Setup logging to view recompiles torch._logging.set_logs(recompilesTrue) for _ in range(5): fn()正如脚本注释所总结的此时会因param_groups[0]中lr的 guard 失败而多次重编译优化器。torch._logging.set_logs是TORCH_LOGS日志体系的 Python API用于观察torch.compile各阶段Dynamo 追踪、图、融合决策、重编译、生成的代码等更完整的用法可参考 torch_logs.py。常见问题与注意事项设备限制torch.compile的 CUDA 后端要求 compute capability 7.0脚本在入口处做设备检查并安全退出CPU 上torch.compile的加速效果与适用性需另行评估首次编译开销第一次调用fn()会触发完整的编译流水线Dynamo 追踪 → 图优化 → Inductor 代码生成 → kernel 编译耗时明显务必通过 warmup 把它排除在基准测量之外重编译陷阱凡是会在迭代中变化并参与计算的非张量标量如普通浮点lr都可能引发 guard 失败与反复重编译应尽量包装为 Tensorfullgraph的选择优化器step()的图结构稳定使用默认fullgraphFalse即可获得完整编译收益同时保留对未知 Python 代码的容错若想强制零图断裂可改用fullgraphTrue出现图断裂时它会直接抛错便于排查结果的机器相关性加速比随硬件、驱动与模型规模变化示例数值约 747 us vs 392 us仅供量级参考应在自己的环境上重新测量并保留assert eager_runtime compiled_runtime作为有效性门禁模型选择不影响结论优化器耗时是参数数量的函数与模型结构无关因此本文的基准方法论可以迁移到任意规模模型。总结本文完整复现并深化了 compiling_optimizer.rst 的核心内容优化器因逐参数更新而天然成为大模型训练的性能瓶颈通过torch.compile(fullgraphFalse)包装opt.step()由 TorchInductor 将多次内存密集的 pointwise 更新融合为更少 kernel可在 GPU 上获得显著的端到端提速。同时我们掌握了科学的测量方法——用torch.utils.benchmark.Timer.blocked_autorange配合 warmup 消除编译噪声、用assert保证结论有效——以及 LR Scheduler 搭配时把 lr 包成 Tensor 避免重编译的关键实践。这些方法论与配套脚本compiling_optimizer_lr_scheduler.py、benchmark.py、torch_logs.py可直接迁移到你的真实训练循环中。赞分享示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载相关推荐PyTorch Lightning高级特性规模化AI研究的终极指南PyTorch Lightning高级特性规模化AI研究的终极指南 PyTorch Lightning是一个强大的深度学习框架它通过封装PyTorch的复杂人工智能深度学习机器学习预训练分布式训练微调Transformers 训练加速实战用 torch.compile 编译优化训练Transformers 训练加速实战用 torch.compile 编译优化训练 本指南围绕 transformers 仓库中的 docs/source/e人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态PyTorch教程使用torch.compile优化器加速训练性能PyTorch教程使用torch.compile优化器加速训练性能 概述 在深度学习模型训练过程中优化器是更新模型参数的核心组件。当处理大型模型时优化器的示例工程上一篇终极指南web-ifc让浏览器IFC处理如此简单下一篇交叉验证方法论张雪峰.skill 如何从碎片言论中提炼出「真信念」的思维框架创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
阅读完成 · 觉得有帮助?
咨询建站