news 2026/9/25 3:50:16

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
使用 torch.compile 编译优化器(Adam)加速 PyTorch 训练:实战指南
  • 示例工程

【免费下载链接】tutorials

PyTorch 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 为每个参数更新操作单独启动 kernel,kernel 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.0(Volta 及以上)的 CUDA 设备。

模型搭建:只关心参数数量

教程选择了一个由 10 层Linear组成的简单顺序模型。作者特别强调:由于我们只基准测试优化器本身,模型的具体结构无关紧要——优化器的性能只取决于参数数量。

import torch model = torch.nn.Sequential( *[torch.nn.Linear(1024, 1024, False, device="cuda") for _ in range(10)] ) input = torch.rand(1024, device="cuda") output = model(input) output.sum().backward()

关键点说明:

  • torch.nn.Linear(1024, 1024, False, device="cuda")中第三个参数bias=False去掉了偏置项,让每一层只含权重矩阵,参数结构更干净;
  • 先执行一次前向传播,再通过output.sum().backward()完成反向传播,为优化器填充梯度。这一步是必须的——没有梯度,opt.step()无事可做;
  • 10 层 × 1024×1024 权重,共约 1000 万参数,规模适中,足以让优化器更新开销成为可观测的测量对象。

设置并运行优化器基准

设备能力检查

由于torch.compile对设备有硬性要求,教程在入口处就做了防护,在不受支持的设备上干净地退出:

# exit cleanly if we are on a device that doesn't 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(), lr=0.01) @torch.compile(fullgraph=False) def fn(): opt.step()

这里两个细节值得展开:

  • fullgraph=False:这是torch.compile的默认设置,表示允许图断裂(graph break)。TorchDynamo 在追踪时若遇到难以捕获的 Python 代码,会中断编译、退回 eager 执行这部分代码,然后继续编译。fullgraph=True则会在遇到第一个图断裂时直接报错。对本例而言,opt.step()内部逻辑可以被完整捕获,fullgraph=False只是保持默认的容错行为。关于图断裂的深入讨论可参考 torch_compile_tutorial.py;
  • 捕获的边界是 Python 函数:torch.compile是装饰器,作用于任意 Python 函数。编译发生时 TorchDynamo 对fn的字节码进行追踪,捕获其中的 PyTorch 算子序列,交给 TorchInductor 生成融合后的底层 kernel(CUDA 下通常是 Triton kernel),后续调用直接复用编译产物。

基准测试辅助函数与测量

教程使用torch.utils.benchmark提供的Timer与blocked_autorange来获得稳定、可统计的耗时:

# Let's define a helpful benchmarking function: import torch.utils.benchmark as benchmark def benchmark_torch_function_in_microseconds(f, *args, **kwargs): t0 = benchmark.Timer( stmt="f(*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(f"eager runtime: {eager_runtime}us") print(f"compiled runtime: {compiled_runtime}us")

几个容易忽略但至关重要的点:

  1. 必须先 warmup 再测量:torch.compile的第一次调用会触发完整编译流程,耗时远高于后续调用(详见 torch_compile_tutorial.py 中首次编译耗时偏大的演示)。教程用 5 次循环预热,让编译产物缓存就位后再计时;
  2. 分别测量 eager 与 compiled:eager 基线直接测opt.step,编译版本测包装函数fn;
  3. assert eager_runtime > compiled_runtime是"门禁":它确保在编译确实带来加速时才继续,若某个环境(如编译产物异常、测量受干扰)下编译版本反而更慢,程序会在此处显式失败,避免输出误导性结论;
  4. 结果具有机器相关性:正如文档明确提示的"Depending on what machine you are using, your exact results may vary",加速比取决于 GPU 型号、驱动、编译缓存等因素,示例数值仅作量级参考。

示例结果

教程给出的单次参考输出为:

  • Eager runtime:约 747.24 us
  • Compiled 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(), lr=torch.tensor(0.01)) sched = torch.optim.lr_scheduler.LinearLR(opt, total_iters=5) @torch.compile(fullgraph=False) def fn(): opt.step() sched.step() # Warmup runs to compile the function for _ in range(5): fn() print(opt.param_groups[0]["lr"])

这里有一个决定成败的细节:把学习率包装成torch.Tensor(lr=torch.tensor(0.01))。

原因是torch.compile会在每次调用时用 guard 校验输入状态是否与已编译版本一致。如果lr是普通 Python 浮点数,调度器每次step()都会修改其值,触发guard 失败,导致函数在每一次迭代都重新编译——编译时间反而成为新的瓶颈。将lr包成 Tensor 后,其值变化以张量数据的形式参与计算,不再触发 guard 层面的重编译。

该脚本还专门演示了如何用日志验证这一点:在非 Tensor 场景下开启重编译日志,就能观察到调度器步进导致的重复编译:

# Setup logging to view recompiles torch._logging.set_logs(recompiles=True) 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 失败与反复重编译,应尽量包装为 Tensor;
  • fullgraph的选择:优化器step()的图结构稳定,使用默认fullgraph=False即可获得完整编译收益,同时保留对未知 Python 代码的容错;若想强制"零图断裂",可改用fullgraph=True,出现图断裂时它会直接抛错便于排查;
  • 结果的机器相关性:加速比随硬件、驱动与模型规模变化,示例数值(约 747 us vs 392 us)仅供量级参考,应在自己的环境上重新测量,并保留assert eager_runtime > compiled_runtime作为有效性门禁;
  • 模型选择不影响结论:优化器耗时是参数数量的函数,与模型结构无关,因此本文的基准方法论可以迁移到任意规模模型。

总结

本文完整复现并深化了 compiling_optimizer.rst 的核心内容:优化器因逐参数更新而天然成为大模型训练的性能瓶颈,通过torch.compile(fullgraph=False)包装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)可直接迁移到你的真实训练循环中。

  • 示例工程

【免费下载链接】tutorials

PyTorch tutorials.

项目地址:https://gitcode.com/gh_mirrors/tuto/tutorials
点击查看免费下载

相关推荐

上一篇:终极指南:web-ifc让浏览器IFC处理如此简单!
下一篇:交叉验证方法论:张雪峰.skill 如何从碎片言论中提炼出「真信念」的思维框架

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 3:49:32

AHK中文编辑器整合版详解:编码配置、热键脚本与一键编译

简介&#xff1a;面向AutoHotkey爱好者和脚本开发者的中文专用编辑器整合包&#xff0c;将SciTE 2.1.0cn中文版、语法高亮、代码折叠、自动完成、括号匹配、查找替换、宏录制、调试以及热更新等常用编辑功能集中到一个轻量环境中&#xff0c;方便中文用户直接上手编写与维护AHK…

作者头像 李华
网站建设 2026/9/25 3:49:27

抖音视频去水印批量下载实战:开源工具douyin-downloader配置与踩坑指南

抖音视频去水印下载这件事&#xff0c;我从2023年就开始折腾了。最开始用的是各种在线解析网站&#xff0c;粘贴链接、点解析、右键保存&#xff0c;一套流程走下来少说半分钟&#xff0c;批量下载更是想都别想。后来陆续试过浏览器插件、手机App、甚至自己抓包写脚本&#xff…

作者头像 李华
网站建设 2026/9/25 3:48:42

STM32H7通过FSMC驱动AD7606,采样率从100k翻倍到200k

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 3:45:01

AgentScope 2.0实战:多Agent调用与RAG as Service企业级落地

最近在中文开发者社区里&#xff0c;AgentScope 这个系统可以说是刷屏级别的存在。AgentScope 2.0、AgentScope Java、RAG as Service、多Agent调用这几个热词&#xff0c;几乎每隔几天就会冒出来一篇新文章&#xff0c;我在好几个技术社群里都被问到过同一个问题&#xff1a;这…

作者头像 李华