Compiled Autograd:为 torch.compile 捕获更大的反向传播图#
创建日期:2024年10月09日 | 最后更新:2026年3月31日 | 最后验证:2024年10月09日
作者: Simon Fan
编译自动求导(Compiled Autograd)如何与
torch.compile交互如何使用 Compiled Autograd API
如何使用
TORCH_LOGS检查日志
PyTorch 2.4
通读 PyTorch 2.x 入门 中的 TorchDynamo 和 AOTAutograd 部分
概述#
Compiled Autograd 是 PyTorch 2.4 中引入的一项 torch.compile 扩展,它允许捕获更大的反向传播图。
虽然 torch.compile 确实会捕获反向传播图,但它是部分捕获的。AOTAutograd 组件预先(ahead-of-time)捕获反向传播图,并存在一定的限制:
前向传播中的图中断(Graph breaks)会导致反向传播中的图中断
Compiled Autograd 通过直接与自动求导引擎集成解决了这些限制,使其能够在运行时捕获完整的反向传播图。具有上述两个特征的模型应尝试使用 Compiled Autograd,并可能获得更好的性能。
然而,Compiled Autograd 也引入了自身的限制:
在反向传播开始时增加了用于缓存查找的运行时开销
由于捕获范围更大,更容易在 Dynamo 中引发重新编译和图中断
注意
Compiled Autograd 处于活跃开发阶段,尚未与所有现有的 PyTorch 功能兼容。关于特定功能的最新状态,请参阅 Compiled Autograd 登陆页面。
设置#
在本教程中,我们将以这个简单的神经网络模型为例。它接收一个 10 维输入向量,通过单层线性层进行处理,并输出另一个 10 维向量。
import torch
class Model(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(10, 10)
def forward(self, x):
return self.linear(x)
基本用法#
在调用 torch.compile API 之前,请确保将 torch._dynamo.config.compiled_autograd 设置为 True。
model = Model()
x = torch.randn(10)
torch._dynamo.config.compiled_autograd = True
@torch.compile
def train(model, x):
loss = model(x).sum()
loss.backward()
train(model, x)
在上述代码中,我们创建了 Model 类的一个实例,并使用 torch.randn(10) 生成了一个 10 维的随机张量 x。我们定义了训练循环函数 train,并使用 @torch.compile 对其进行装饰以优化执行。当调用 train(model, x) 时:
Python 解释器调用 Dynamo,因为此调用被
@torch.compile装饰。Dynamo 拦截 Python 字节码,模拟其执行并将操作记录到图中。
AOTDispatcher禁用钩子并调用自动求导引擎,计算model.linear.weight和model.linear.bias的梯度,并将操作记录到图中。使用torch.autograd.Function,AOTDispatcher 重写了train的前向和反向传播实现。Inductor 生成一个对应于 AOTDispatcher 前向和反向传播优化实现的函数。
Dynamo 将优化后的函数设置为下一个由 Python 解释器评估的函数。
Python 解释器执行优化后的函数,该函数执行
loss = model(x).sum()。Python 解释器执行
loss.backward(),调用自动求导引擎,由于我们设置了torch._dynamo.config.compiled_autograd = True,该调用路由至 Compiled Autograd 引擎。Compiled Autograd 计算
model.linear.weight和model.linear.bias的梯度,并将操作记录到图中,包括它遇到的任何钩子。在此过程中,它会记录之前由 AOTDispatcher 重写的反向传播过程。随后,Compiled Autograd 生成一个新函数,该函数对应于loss.backward()的完全追踪实现,并以推理模式运行torch.compile来执行它。同样的步骤递归地应用于 Compiled Autograd 图,但这次 AOTDispatcher 不需要划分图。
检查 Compiled Autograd 日志#
使用 TORCH_LOGS 环境变量运行脚本。
若只想打印 Compiled Autograd 图,使用
TORCH_LOGS="compiled_autograd" python example.py。若要以牺牲性能为代价,打印包含更多张量元数据和重新编译原因的图,使用
TORCH_LOGS="compiled_autograd_verbose" python example.py。
重新运行上述代码片段,Compiled Autograd 图应已记录到 stderr 中。某些图节点名称将带有 aot0_ 前缀,这些对应于之前在 AOTAutograd 反向传播图 0 中预编译的节点,例如 aot0_view_2 对应于 id=0 的 AOT 反向传播图中的 view_2。
在下图中,红色框中封装的是没有 Compiled Autograd 时由 torch.compile 捕获的 AOT 反向传播图。
注意
这是我们将要在其上调用 torch.compile 的图,而不是优化后的图。Compiled Autograd 本质上生成了一些未优化的 Python 代码来表示整个 C++ 自动求导执行过程。
使用不同的标志编译前向和反向传播#
你可以为两次编译使用不同的编译器配置,例如,即使前向传播中有图中断,反向传播也可以是 fullgraph(全图)。
def train(model, x):
model = torch.compile(model)
loss = model(x).sum()
torch._dynamo.config.compiled_autograd = True
torch.compile(lambda: loss.backward(), fullgraph=True)()
或者你可以使用上下文管理器,它将应用于其范围内的所有自动求导调用。
def train(model, x):
model = torch.compile(model)
loss = model(x).sum()
with torch._dynamo.compiled_autograd.enable(torch.compile(fullgraph=True)):
loss.backward()
Compiled Autograd 解决了 AOTAutograd 的某些限制#
前向传播中的图中断不再必然导致反向传播中的图中断。
@torch.compile(backend="aot_eager")
def fn(x):
# 1st graph
temp = x + 10
torch._dynamo.graph_break()
# 2nd graph
temp = temp + 10
torch._dynamo.graph_break()
# 3rd graph
return temp.sum()
x = torch.randn(10, 10, requires_grad=True)
torch._dynamo.utils.counters.clear()
loss = fn(x)
# 1. base torch.compile
loss.backward(retain_graph=True)
assert(torch._dynamo.utils.counters["stats"]["unique_graphs"] == 3)
torch._dynamo.utils.counters.clear()
# 2. torch.compile with compiled autograd
with torch._dynamo.compiled_autograd.enable(torch.compile(backend="aot_eager")):
loss.backward()
# single graph for the backward
assert(torch._dynamo.utils.counters["stats"]["unique_graphs"] == 1)
在第一个 torch.compile 示例中,我们可以看到由于编译函数 fn 中存在 2 次图中断,产生了 3 个反向传播图。而在第二个带有 Compiled Autograd 的 torch.compile 示例中,尽管存在图中断,我们仍看到了一个完整捕获的反向传播图。
注意
在追踪 Compiled Autograd 捕获的反向传播钩子时,Dynamo 仍然可能发生图中断。
反向传播钩子现在可以被捕获。
@torch.compile(backend="aot_eager")
def fn(x):
return x.sum()
x = torch.randn(10, 10, requires_grad=True)
x.register_hook(lambda grad: grad+10)
loss = fn(x)
with torch._dynamo.compiled_autograd.enable(torch.compile(backend="aot_eager")):
loss.backward()
图中应该存在一个 call_hook 节点,Dynamo 稍后会将其内联到以下内容中:
Compiled Autograd 的常见重新编译原因#
由于损失值(loss)的自动求导结构发生变化
torch._dynamo.config.compiled_autograd = True
x = torch.randn(10, requires_grad=True)
for op in [torch.add, torch.sub, torch.mul, torch.div]:
loss = op(x, x).sum()
torch.compile(lambda: loss.backward(), backend="eager")()
在上面的示例中,我们在每次迭代中调用不同的算子,导致 loss 每次都跟踪不同的自动求导历史。你应该会看到一些重新编译消息:Cache miss due to new autograd node(由于新的自动求导节点导致缓存未命中)。
由于张量形状改变
torch._dynamo.config.compiled_autograd = True
for i in [10, 100, 10]:
x = torch.randn(i, i, requires_grad=True)
loss = x.sum()
torch.compile(lambda: loss.backward(), backend="eager")()
在上面的示例中,x 的形状发生了变化,Compiled Autograd 会在第一次变化后将 x 标记为动态形状张量。你应该会看到重新编译消息:Cache miss due to changed shapes(由于形状改变导致缓存未命中)。
结论#
在本教程中,我们介绍了使用 Compiled Autograd 的 torch.compile 生态系统概述、Compiled Autograd 的基础知识以及一些常见的重新编译原因。请关注 dev-discuss 上的深度解析文章。