PyTorch torch.compile 实战 FAQ:训练、分布式、性能诊断与 NumPy 编译全指南
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
本文以 PyTorch 官方 FAQ 文档 docs/source/user_guide/torch_compiler/torch.compiler_faq.md 为骨架,结合本仓库中 TorchDynamo、AOTAutograd、TorchInductor 与 NumPy 桥接的真实源码,系统回答torch.compile使用中最高频的十二类问题:训练与分布式支持、整图导出、崩溃与精度排查、编译耗时、加速原理、Graph Break、OOM、torch.func组合以及 NumPy 代码编译。读完本文,你将掌握一套从"能跑"到"跑得快、跑得稳、能排查"的完整实战方法。
torch.compile 支持训练吗?
支持。torch.compile通过 AOTAutograd 捕获反向图来完成训练支持,其工作流程分为五步:
.forward()图与optimizer.step()由 TorchDynamo 的 Pythonevalframe前端捕获;- 对 TorchDynamo 捕获的每一段 forward 图,使用 AOTAutograd 生成对应的 backward 图段;
- 每对 forward/backward 图(可选地)经过 min-cut 分区,在 forward 与 backward 之间仅保存最少的中间状态(内存最小化);
- forward/backward 对被包装进
autograd.function模块; - 用户代码调用
.backward()时,仍然触发 eager 的 autograd 引擎,每个已编译的 backward 图被当作一个算子执行,同时运行任何未编译 eager 算子的.backward()。
支持分布式代码吗?
torch.compile支持DistributedDataParallel(DDP),其他分布式训练库的支持正在评估中。
分布式代码对 Dynamo 的挑战在于:AOTAutograd 会同时展开 forward 与 backward,并向后端提供两张图,而分布式场景希望通信与计算尽量重叠。Eager PyTorch 通过 autograd hooks、module hooks 以及对模块状态的修改/变更,以不同方式为 DDP/FSDP 实现这一目标;而在 Dynamo 朴素应用下,本应在 backward 某算子后立即执行的 hook,可能因 AOTAutograd 编译函数与 dispatcher hooks 的交互方式,被延迟到整个已编译 backward 区域结束后才执行。
针对 DDP 的优化策略在 torch/_dynamo/backends/distributed.py 中实现。该文件的核心组件是DDPOptimizer类(distributed.py 中 class DDPOptimizer),主要思路是在DDP bucket 边界上产生 graph break:DDP 中每个节点需要与其他节点同步权重时,会把梯度和参数组织成 bucket,从而减少通信次数,并允许节点向其他等待节点广播其部分梯度。从源码结构看,模块中还定义了Bucket数据结构(记录params、nodes、param_ids等)与SubmodCompiler,用于把切分后的子图逐个编译。
分布式代码中的 graph break 意味着:Dynamo 及其后端只能优化分布式程序的计算开销,而无法优化其通信开销。Graph break 可能干扰编译加速——图变小会让编译器失去融合机会——但大多数当前计算优化属于局部融合,图增大带来的收益存在边际递减,因此该方案在实践中通常是够用的。
还需要导出整图(full graph)吗?
对绝大多数模型而言不需要,直接使用torch.compile()即可。但以下少数场景需要整图,可通过torch.compile(..., fullgraph=True)强制保证:
- 大规模训练任务:例如需要 pipeline parallelism 和其他高级分片策略的超大规模训练(原文以 $250K+ 量级的训练运行为例);
- 推理优化器:如 TensorRT、AITemplate 这类比训练优化器更激进地依赖融合的推理加速方案;
- 移动端训练或推理。
未来工作将包括:把通信操作追踪进图中、协调通信与计算优化、以及直接优化通信操作本身。
为什么我的代码崩溃了?
如果代码在未启用torch.compile时一切正常、启用后才崩溃,第一步是确定失败发生在技术栈的哪一层。按以下顺序排查,只有上一步通过才尝试下一步:
torch.compile(..., backend="eager"):仅运行 TorchDynamo 的 forward 图捕获,然后用 PyTorch 运行捕获的图。若失败,问题出在TorchDynamo;torch.compile(..., backend="aot_eager"):TorchDynamo 捕获 forward 图后,由 AOTAutograd 追踪 backward 图,不再进行额外后端编译步骤,forward/backward 图均由 PyTorch eager 执行。若失败,问题出在AOTAutograd;torch.compile(..., backend="inductor"):TorchDynamo 捕获 forward 图,AOTAutograd 追踪 backward 图,再交给 TorchInductor 编译。若失败,问题出在TorchInductor。
这套"逐步下沉"的二分定位法能快速把故障隔离到 Dynamo 前端、AOTAutograd 还是 Inductor 代码生成层。
为什么编译很慢?
Dynamo 编译阶段
TorchDynamo 内置统计函数,用于收集并展示每个编译阶段耗时。执行torch._dynamo后调用torch._dynamo.utils.compile_times()即可获取。其实现位于 torch/_dynamo/utils.py#L923-L959:默认repr="str"时返回"按函数名统计的 TorchDynamo 编译时间"的可打印字符串(使用tabulate输出表格);也支持repr="csv"返回表头与行数据以便记录日志;aggregate=True会把多次编译(如切分后的图)的值累加为单个值。此外模块注册了atexit钩子dump_compile_times(),进程退出时会自动打印聚合后的编译时间。
Inductor 编译阶段
TorchInductor 内置统计与 trace 功能,可展示各编译阶段耗时、输出代码、输出图可视化以及 IR dump。运行方式:
env TORCH_COMPILE_DEBUG=1 python repro.py其配置位于 torch/_inductor/config.py#L2933-L2999 的class trace,注意:
enabled由环境变量TORCH_COMPILE_DEBUG控制(默认"0");fx_graph/fx_graph_transformed:保存分解后、变换前/后的 FX 图;ir_pre_fusion/ir_post_fusion:保存融合前/后的 TorchInductor IR;output_code:把生成的代码复制到 trace 目录;graph_diagram:输出融合后图的可视化(默认禁用,因生成代价高),可通过INDUCTOR_POST_FUSION_SVG或INDUCTOR_POST_FUSION_GRAPH环境变量开启,格式由TORCH_COMPILE_GRAPH_FORMAT控制,默认 svg;draw_orig_fx_graph:输出原始 FX 图可视化(同样默认禁用)。
debug trace 目录中的每个文件可通过torch._inductor.config.trace.*单独开启/关闭。
过度重编译(Excessive Recompilation)
TorchDynamo 编译一个函数(或其中一部分)时,会基于对 locals 与 globals 的假设做编译优化,并将这些假设表达为运行时检查的 guards。一旦某 guard 失败,Dynamo 会重编译该函数,上限为torch._dynamo.config.recompile_limit次。相关配置见 torch/_dynamo/config.py#L121-L149:默认recompile_limit = 8,另有accumulated_recompile_limit = 256限制累积重编译次数,以及fail_on_recompile_limit_hit = False(设为 True 可在命中上限时直接报错而非继续)。
若程序命中缓存上限,首先需要确定是哪个 guard 失败、程序中哪部分触发。可使用TORCH_TRACE/tlparse或TORCH_LOGS=recompiles追踪问题根源,详见仓库中 torch.compiler_troubleshooting 相关内容。
为什么在生产环境还在重编译?
某些场景下(例如延迟敏感的线上流量服务),你可能不希望程序热身之后还出现意外编译。TorchDynamo 为此提供了另一种模式:复用此前已编译的图,但不再生成新图:
frozen_toy_example = dynamo.run(toy_example) frozen_toy_example(torch.randn(10), torch.randn(10))你们是如何加速我的代码的?
加速 PyTorch 代码主要有三大类手段:
内核融合(Kernel Fusion)
- 垂直融合:融合顺序操作以减少过多的读写。例如融合两次连续的 cos,可将 2 次读 + 2 次写降为 1 次读 + 1 次写;
- 水平融合:最简单的例子是 batch 场景下单个矩阵与一批样本相乘,更一般的场景是 grouped GEMM——将一组矩阵乘法统一调度执行。
乱序执行(Out of Order Execution):编译器通过提前查看图内精确的数据依赖关系,决定执行某个节点最合适的时机,以及哪些缓冲区可以被复用。
自动任务放置(Automatic Work Placement):与乱序执行类似,但更进一步——通过将图的节点匹配到物理硬件或内存等资源,设计出合适的执行调度。
以上是加速 PyTorch 代码的通用原则,不同后端会在"优化什么"上做不同取舍。例如 Inductor 先尽可能融合,然后才生成 Triton 内核。Triton 额外带来加速的原因包括:每个 Streaming Multiprocessor 内的自动内存合并(memory coalescing)、内存管理、调度,以及为分块(tiled)计算而设计。
无论使用哪个后端,都建议采用benchmark-and-see 方法:使用 PyTorch profiler,目视检查生成的内核,亲自观察实际发生了什么。
为什么我没有看到加速?
Graph Break
看不到预期加速的主要原因是过多的 graph break。什么是 graph break?
def some_fun(x): ... torch.compile(some_fun)(x) ...TorchDynamo 会尝试把some_fun()中所有 torch/tensor 操作编译进单个 FX 图,但也可能无法把全部内容捕获进一张图。有些 graph break 的原因对 TorchDynamo 而言是无法逾越的——例如调用 PyTorch 之外的 C 扩展对 TorchDynamo 不可见,它可能做任意事情,而 TorchDynamo 无法引入必要的 guards 来保证编译后的程序可以安全复用。
为了最大化性能,graph break 越少越好。
定位 graph break 的原因
使用torch._dynamo.explain可识别程序中所有 graph break 及其原因。该工具在指定函数上运行 TorchDynamo 并聚合遇到的 graph break,其实现位于 torch/_dynamo/eval_frame.py#L1932-L1981(内部通过optimize累积 graph、break_reasons、op_count、ops_per_graph,并导出 guards)。示例:
import torch import torch._dynamo as dynamo def toy_example(a, b): x = a / (torch.abs(a) + 1) print("woo") if b.sum() < 0: b = b * -1 return x * b explanation = dynamo.explain(toy_example)(torch.randn(10), torch.randn(10)) print(explanation) """ Graph Count: 3 Graph Break Count: 2 Op Count: 5 Break Reasons: Break Reason 1: Reason: builtin: print [<class 'torch._dynamo.variables.constant.ConstantVariable'>] False User Stack: <FrameSummary file foo.py, line 5 in toy_example> Break Reason 2: Reason: generic_jump TensorVariable() User Stack: <FrameSummary file foo.py, line 6 in torch_dynamo_resume_in_toy_example_at_5> Ops per Graph: ... Out Guards: ... """若希望在遇到第一个 graph break 时直接报错,可以禁用 python fallback,使用fullgraph=True——熟悉基于 export 的编译器的话,这个选项不会陌生:
def toy_example(a, b): ... torch.compile(toy_example, fullgraph=True, backend=<compiler>)(a, b)为什么我改了代码却没有重编译?
如果通过env TORCHDYNAMO_DYNAMIC_SHAPES=1 python model.py开启了动态形状,那么形状变化时代码不会重编译。动态形状支持避免了"形状变化小于 2 倍"情况下的重编译,这在图像尺寸多变的 CV 场景或序列长度可变的 NLP 场景尤其有用;推理场景中,batch size 事先往往无法确定(取决于不同客户端应用的实际请求),动态形状也很有价值。
一般而言,TorchDynamo 会尽力避免无谓的重编译:例如 TorchDynamo 发现 3 张图,而你的改动只影响了其中 1 张,那么只有那张图会重编译。因此避免编译缓慢的另一个技巧是热身(warmup):先编译一次模型,之后的编译会快得多。冷启动编译时间是官方持续跟踪的指标。
为什么结果不正确?
精度问题可以借助环境变量TORCHDYNAMO_REPRO_LEVEL=4做最小化复现,其工作方式类似git bisect。一个完整的复现示例:
TORCHDYNAMO_REPRO_AFTER="aot" TORCHDYNAMO_REPRO_LEVEL=4之所以需要它,是因为下游编译器(无论是 Triton 代码还是 C++ 后端)会做代码生成,其数值结果可能与 eager 存在细微差异,却对训练稳定性产生重大影响。因此该精度调试器对检测 codegen 或后端编译器中的 bug 非常有用。
如果希望确保 torch 与 triton 之间的随机数生成一致,可以开启:
torch._inductor.config.fallback_random = True该配置项默认值为False,定义于 torch/_inductor/config.py#L965。
为什么出现 OOM?
Dynamo 仍是 alpha 阶段产品,OOM 有若干来源。若遇到 OOM,按以下顺序尝试关闭配置,并到 GitHub 提 issue 以便从根源上解决:
- 动态形状:若在使用动态形状,尝试关闭(默认已关闭):
env TORCHDYNAMO_DYNAMIC_SHAPES=0 python model.py - CUDA graphs:Inductor 中默认启用 Triton 的 CUDA graphs,移除可能缓解部分 OOM 问题:
torch._inductor.config.triton.cudagraphs = False从 torch/_inductor/config.py#L1978-L1979 的源码看,
cudagraphs也支持通过环境变量TORCHINDUCTOR_CUDAGRAPHS=1开启,其背后还有 cudagraph trees 内存池机制(cudagraph_trees)。
torch.func 能与 torch.compile 组合使用吗(grad/vmap 变换)?
对使用了torch.compile的函数应用torch.func变换是可行的:
import torch @torch.compile def f(x): return torch.sin(x) def g(x): return torch.grad(f)(x) x = torch.randn(2, 3) g(x)在 torch.compile 处理的函数内部调用 torch.func 变换
编译torch.func.grad:
import torch def wrapper_fn(x): return torch.func.grad(lambda x: x.sin().sum())(x) x = torch.randn(3, 3, 3) grad_x = torch.compile(wrapper_fn)(x)编译torch.vmap:
import torch def my_fn(x): return torch.vmap(lambda x: x.sum(1))(x) x = torch.randn(3, 3, 3) output = torch.compile(my_fn)(x)编译不受支持函数(逃生通道)
对于其他变换,可使用torch._dynamo.allow_in_graph作为 workaround。allow_in_graph是一个逃生通道:如果代码无法与torch.compile(它内省 Python 字节码)协作,但你相信它可以通过符号追踪方式(类似jax.jit)工作,就使用allow_in_graph。其实现见 torch/_dynamo/decorators.py#L199-L222——它跳过对函数的符号内省,直接把函数写入图中(注册进trace_rules._allowed_callable_ids)。
使用allow_in_graph标注函数时,必须满足以下要求:
- 函数的所有输出只依赖输入,不依赖任何被捕获的 Tensor;
- 函数是函数式的:不改变任何状态。这一点可以放宽——我们实际上支持"从外部看是函数式"的函数:可以有就地 PyTorch 操作,但不能改变全局状态或函数输入;
- 函数不抛出依赖数据的错误(data-dependent errors)。
import torch @torch.compile def f(x): return torch._dynamo.allow_in_graph(torch.vmap(torch.sum))(x) x = torch.randn(2, 3) f(x)一个常见陷阱:用allow_in_graph标注调用了nn.Module的函数。因为此时输出依赖nn.Module的参数,无法满足"只依赖输入"。解决方法是使用torch.func.functional_call提取模块状态。
NumPy 能与 torch.compile 协作吗?
从 PyTorch 2.1 开始,torch.compile能理解:纯 NumPy 程序(操作 NumPy 数组),以及通过x.numpy()、torch.from_numpy等函数在 PyTorch 与 NumPy 间转换的混合程序。
torch.compile 支持哪些 NumPy 特性?
torch.compile内的 NumPy 遵循 NumPy 2.0 预发布版语义。总体而言,torch.compile能追踪大多数 NumPy 构造;无法追踪时回退到 eager,让 NumPy 执行该段代码。即便如此,仍有少数特性的语义与 NumPy 略有偏差:
- NumPy 标量:建模为 0-D 数组,即
np.float32(3)在torch.compile下返回 0-D 数组。为避免 graph break,最好直接使用该 0-D 数组;若破坏代码,可将 NumPy 标量转换为对应 Python 标量类型bool/int/float作为 workaround; - 负步长:
np.flip与负步长切片返回副本; - 类型提升:NumPy 2.0 中类型提升规则会改变(见 NEP 50)。
torch.compile实现的是 NEP 50 而非当前即将废弃的旧规则; {tril,triu}_indices_from/{tril,triu}_indices:返回数组而非数组元组。
以下特性不支持追踪,会优雅回退到 NumPy 执行:
- 非数值 dtype:datetime、string、char、void、structured dtypes 与 recarrays;
- 长 dtype
np.float128/np.complex256与部分无符号 dtypenp.uint16/np.uint32/np.uint64; ndarray子类;- 掩码数组(masked arrays);
- 深奥的 ufunc 机制,如
axes=[(n,k),(k,m)->(n,m)]与 ufunc 方法(如np.add.reduce); - 对
complex64/complex128数组排序/排序相关操作; - NumPy
np.poly1d与np.polynomial; - 两个及以上返回值函数中的位置参数
out1, out2(out=tuple可用); __array_function__、__array_interface__、__array_wrap__;ndarray.ctypes属性。
可以用 torch.compile 编译 NumPy 代码吗?
当然可以。torch.compile原生理解 NumPy 代码,把它当作 PyTorch 代码处理,只需用torch.compile装饰器包裹:
import torch import numpy as np @torch.compile def numpy_fn(X: np.ndarray, Y: np.ndarray) -> np.ndarray: return np.sum(X[:, :, None] * Y[:, None, :], axis=(-2, -1)) X = np.random.randn(1024, 64) Y = np.random.randn(1024, 64) Z = numpy_fn(X, Y) assert isinstance(Z, np.ndarray)用环境变量TORCH_LOGS=output_code执行此示例,可以看到torch.compile把乘法与求和融合进了单个 C++ 内核,并通过 OpenMP 并行执行(原生 NumPy 是单线程的)。这很容易让 NumPy 代码提速n倍——n即处理器的核心数。追踪 NumPy 代码同样支持编译代码内部的 graph break。
能通过 torch.compile 在 CUDA 上执行 NumPy 代码并计算梯度吗?
可以。只需在torch.device("cuda")上下文内执行代码:
import torch import numpy as np @torch.compile def numpy_fn(X: np.ndarray, Y: np.ndarray) -> np.ndarray: return np.sum(X[:, :, None] * Y[:, None, :], axis=(-2, -1)) X = np.random.randn(1024, 64) Y = np.random.randn(1024, 64) with torch.device("cuda"): Z = numpy_fn(X, Y) assert isinstance(Z, np.ndarray)此时numpy_fn会在 CUDA 上执行:torch.compile自动把X、Y从 CPU 搬到 CUDA,再把结果Z从 CUDA 搬回 CPU。若同一程序中多次执行该函数,应避免这些昂贵的拷贝——只需让numpy_fn接受 CUDA Tensor 并返回 Tensor,即使用torch.compiler.wrap_numpy(实现见 torch/compiler/init.py#L546-L575,将"np.ndarray → np.ndarray"的函数变为"torch.Tensor → torch.Tensor"的函数,与torch.compile(fullgraph=True)配合使用):
@torch.compile(fullgraph=True) @torch.compiler.wrap_numpy def numpy_fn(X, Y): return np.sum(X[:, :, None] * Y[:, None, :], axis=(-2, -1)) X = torch.randn(1024, 64, device="cuda") Y = torch.randn(1024, 64, device="cuda") Z = numpy_fn(X, Y) assert isinstance(Z, torch.Tensor) assert Z.device.type == "cuda"这里我们显式地在 CUDA 内存中创建 tensor 并传入函数,所有计算都在 CUDA 设备上进行。wrap_numpy负责在torch.compile层面把任何torch.Tensor输入标记为具有np.ndarray语义的输入。在编译器内部标记 tensor 是非常廉价的操作,运行时不会发生数据拷贝或移动。
使用该装饰器还可以对 NumPy 代码求导:
@torch.compile(fullgraph=True) @torch.compiler.wrap_numpy def numpy_fn(X, Y): return np.mean(np.sum(X[:, :, None] * Y[:, None, :], axis=(-2, -1))) X = torch.randn(1024, 64, device="cuda", requires_grad=True) Y = torch.randn(1024, 64, device="cuda") Z = numpy_fn(X, Y) assert isinstance(Z, torch.Tensor) Z.backward() # X.grad now holds the gradient of the computation print(X.grad)之所以使用fullgraph=True,是因为 graph break 在此场景中有问题:发生 graph break 时,需要物化 NumPy 数组;而 NumPy 数组没有device或requires_grad的概念,graph break 时会丢失这些信息。我们无法跨 graph break 传播梯度,因为 break 处的代码可能执行任意无法求导的代码。另一方面,在 CUDA 执行场景中,可以像第一个示例那样用torch.device("cuda")上下文管理器绕过:
@torch.compile @torch.compiler.wrap_numpy def numpy_fn(X, Y): prod = X[:, :, None] * Y[:, None, :] print("oops, a graph break!") return np.sum(prod, axis=(-2, -1)) X = torch.randn(1024, 64, device="cuda") Y = torch.randn(1024, 64, device="cuda") with torch.device("cuda"): Z = numpy_fn(X, Y) assert isinstance(Z, torch.Tensor) assert Z.device.type == "cuda"在 graph break 期间,中间 tensor 仍需要移动到 CPU,但 break 之后恢复追踪时,图的其余部分仍在 CUDA 上追踪。由于这种 CUDA ↔ CPU 与 CPU ↔ CUDA 的移动,graph break 在 NumPy 场景下代价相当高,应当避免;但至少它们允许追踪复杂代码片段。
如何调试 torch.compile 下的 NumPy 代码?
调试 JIT 编译代码颇具挑战性——现代编译器复杂,报错令人望而生畏。torch.compiler_troubleshooting文档包含一些排查技巧。若仍无法定位问题,还有几个 NumPy 专属工具:
- 判断 bug 是否完全在 PyTorch 代码中:禁用对 NumPy 函数的追踪:
from torch._dynamo import config config.trace_numpy = False - 判断 bug 是否在被追踪的 NumPy 代码中:不用
torch.compile,以 PyTorch 为后端直接 eager 执行 NumPy 代码——导入import torch._numpy as np。torch._numpy是用 PyTorch 实现的 NumPy 的 Python 实现,被torch.compile内部用来把 NumPy 代码转换为 PyTorch 代码。它仅用于调试,绝不是 PyTorch API 的替代品:性能差得多,且作为私有 API可能随时变更。它相当易读、易改,若发现 bug 欢迎提交 PR 修复或直接提 issue。 - 若程序在导入
torch._numpy as np后能正常工作,那么 bug 很可能在 TorchDynamo 中。此时请附带一个最小复现(minimal reproducer)提 issue。
编译了 NumPy 代码却没看到加速?
最佳起点是本文"为什么我没有看到加速"一节。部分 graph break 可能源于使用了不支持的特性(见上文"支持哪些 NumPy 特性"清单)。更一般地,有些广泛使用的 NumPy 特性与编译器配合不佳:
- 就地修改(in-place):在编译器内部难以推理,通常比非就地版本性能更差,应尽量避免;
out=参数:同样建议避免,改用非就地操作并让torch.compile优化内存使用;- 数据依赖操作:如通过布尔掩码的掩码索引(masked indexing),以及数据依赖的控制流(
if/while构造)。
细粒度追踪该用哪个 API?
有时需要把代码的一小部分排除在torch.compile编译之外。以下是最常用的答案(更多信息见 TorchDynamo 细粒度追踪文档)。
如何让某个函数 graph break?
"在函数上 graph break"不足以精确表达你想要的行为,需要明确你的具体用例:
- 想禁用该函数帧及其递归调用帧的编译:用
torch._dynamo.disable; - 想让某个算子(如
fbgemm)使用 eager 模式:用torch._dynamo.disallow_in_graph。
较少见的用例:
- 想只禁用该函数帧、但恢复其递归调用帧上的 TorchDynamo:
torch._dynamo.disable(recursive=False); - 想阻止某个函数帧被内联:在该函数开头使用
torch._dynamo.graph_break。
torch._dynamo.disable 与 torch._dynamo.disallow_in_graph 的区别
disallow_in_graph工作在算子层面——更具体地说,是你能在 TorchDynamo 提取的图中看到的算子(对应实现见 torch/_dynamo/decorators.py#L774)。disable工作在函数帧层面,决定 TorchDynamo 是否要查看该函数帧(实现见 torch/_dynamo/decorators.py#L90-L128,recursive=True时用DisableContext完全跳过该帧及其递归调用的帧;recursive=False时跳过该函数代码关联的帧,但仍处理递归调用的帧)。
torch._dynamo.disable 与 torch._dynamo.skip 的区别
::: {note}torch._dynamo.skip已弃用。 :::
你大概率只需要torch._dynamo.disable。但在极少数场景下需要更细的控制:假设只想禁用a_fn上的追踪,同时继续在aa_fn与ab_fn中恢复追踪——此时可用torch._dynamo.disable(recursive=False)。旧版本中该功能由torch._dynamo.skip提供,现已由disable的recursive标志支持。
小结
torch.compile的完整技术栈由三层协作构成:TorchDynamo(字节码级前端捕获与 guards)、AOTAutograd(forward/backward 联合捕获与内存最小化)、TorchInductor(IR 优化、融合与 Triton/C++ 代码生成)。遇到问题时,按"backend=eager → aot_eager → inductor"逐层下沉定位故障;追求性能时,用dynamo.explain找出 graph break 并尽量消除;对动态形状、NumPy 桥接、torch.func组合等高级用法,本仓库的 torch/_dynamo、torch/_inductor/config.py、torch/compiler/init.py 都是深入阅读的起点。关于更多排查手段,可继续阅读 torch.compiler_troubleshooting。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考