首页
/ PyTorch torch.compile 中的 torch._dynamo.nonstrict_trace:消除图断点的完整编程模型与源码解析

PyTorch torch.compile 中的 torch._dynamo.nonstrict_trace:消除图断点的完整编程模型与源码解析

2026-09-09 17:06:33作者:蔡怀权

本文围绕 PyTorch 官方文档《Use torch._dynamo.nonstrict_trace》展开,讲解如何在 torch.compile(fullgraph=True) 区域内消除图断点(graph break):何时需要 nonstrict_trace、它的四项前置要求、装饰器与函数调用两种用法,并结合 torch/_dynamo 源码与 test/dynamo/test_decorators.py 测试用例,深入剖析其内部注册机制、与 allow_in_graph 的区别以及常见的报错边界。

1. 问题背景:fullgraph=True 下遇到图断点会直接报错

当你希望编译区域是一整张图(fullgraph=True)时,任何内部触发的图断点都会导致 Dynamo 抛出错误。文档中给出的最小复现场景是:

import torch

def get_magic_num():
    # 这个显式的 graph_break 调用模拟任何一种 Dynamo 图断点,
    # 例如函数用 C 实现,或用到了 Dynamo 尚不支持的 Python 语言特性。
    torch._dynamo.graph_break()
    return torch.tensor([42])

@torch.compile(fullgraph=True)
def func(x):
    n = get_magic_num()
    return x + n

try:
    func(torch.rand(10))
except Exception as e:
    print(e)

运行后 Dynamo 会报错,因为它在 fullgraph=True 的编译区域内看到了一个图断点。

从源码结构看,torch._dynamo.graph_break() 本身是通过 _disallow_in_graph_helper(throw_if_not_allowed=False) 注册到 Dynamo 的禁止入图名单中的(见 decorators.py),它在被追踪到时即强制中断当前图。这正是需要 nonstrict_trace 的出发点:你确定这个函数可以被"非严格追踪"(non-strict tracing),希望把它的函数体展开进最终 FX 图,而不是让它在调用点断开

当图断点发生时,用户通常有几种应对方式(文档给出的完整清单):

  1. 图断点源于 Dynamo 尚不支持的语言特性 —— 重写代码,或向项目提交 issue;
  2. 图断点源于调用了一个 C 实现函数 —— 可以尝试改用自定义算子(custom op),或提供一个 polyfill(Python 参考实现)让 Dynamo 追踪穿过它;
  3. 最坏情况 —— 编译器内部错误,此时通常需要提交 issue。

除了以上选项,PyTorch 还专门提供了 torch._dynamo.nonstrict_trace 作为替代方案。

2. nonstrict_trace 的语义:Dynamo 不进入函数体,AOT Autograd 会进入

先明确它"非严格"的确切含义。torch.compiler.nonstrict_trace 的文档字符串给出了权威解释(见 torch/compiler/init.py):

A nonstrict-traced function appears as an opaque call in the dynamo graph. Dynamo does not trace into the function body (hence the "nonstrict"), but aot_autograd will trace into it.

也就是说:

  • Dynamo 层面:该函数在 Dynamo 图中呈现为一个不透明调用(opaque call),Dynamo 不逐行追踪其函数体,因此函数体内部即使出现 graph_break()、C 实现调用等本会导致断点的内容,也不会中断 Dynamo 的追踪;
  • AOT Autograd 层面:AOT Autograd 会真正追踪进函数体,最终 FX 图将包含该函数内部发生的全部张量操作。

不同 backend 下的行为差异(文档字符串中明确说明):

  • backend="eager":原始 Python 函数直接运行;
  • backend="aot_eager":运行 AOT Autograd 追踪出的图;
  • backend="inductor":追踪出的图被 Inductor 编译。

并且训练是被支持的:可以对输出调用 .backward(),梯度会流经 nonstrict-trace 的函数(源码文档示例中即演示了 out.sum().backward() 的场景)。

3. 使用前提:nonstrict_trace 的四项要求

文档明确列出了调用方能被 nonstrict_trace 处理的必要条件,这四点也是排查报错的核心依据:

  1. 满足通用非严格追踪的全部要求:函数内部只能出现纯张量计算与受支持的 Python 控制流,不能依赖副作用与外部可变状态的交互顺序;
  2. 输入与输出必须是基本类型或已注册类型:即 intfloatlistdicttorch.Tensor 等基础类型,或者是已经注册到 torch.utils._pytree 的用户自定义类型;
  3. 函数必须定义在 torch.compile 区域之外:不能在编译函数内部用 torch._dynamo.nonstrict_trace(trace_me) 去包装一个刚刚在局部作用域定义出来的函数;
  4. 非输入值会被当作常量:函数读到的任何非输入值(例如全局张量、闭包捕获的张量)都被视为常量,不会对它们建 guard

其中第 4 条有两个重要推论,源码 docstring 将其归入"Dangerous patterns":

  • 闭包/全局捕获的张量不会被求梯度 —— 需要梯度时请把张量作为显式参数传入;
  • 编译区域内的其他代码不应依赖该函数的副作用(修改共享状态后,编译图可能重排或消除操作,导致顺序假设被破坏)。

这些边界都有对应的测试用例验证,见第 6 节。

4. 两种用法:装饰器形式与区域内调用形式

4.1 装饰器形式

文档主示例:给 get_magic_num 加上 @torch._dynamo.nonstrict_trace 后,图断点被消除,fullgraph=True 下不再报错:

@torch._dynamo.nonstrict_trace
def get_magic_num():
    # 模拟 Dynamo 图断点:C 实现或不被支持的语言特性
    torch._dynamo.graph_break()
    return torch.tensor([42])

@torch.compile(fullgraph=True)
def func(x):
    n = get_magic_num()
    return x + n

print(func(torch.rand(10)))
# No graph break and no error.

4.2 在 torch.compile 区域内直接调用

文档还指出可以在编译区域内以函数调用的方式使用:

def get_magic_num():
    torch._dynamo.graph_break()
    return torch.tensor([42])

@torch.compile(fullgraph=True)
def func(x):
    n = torch._dynamo.nonstrict_trace(get_magic_num)()
    return x + n

print(func(torch.rand(10)))
# No graph break and no error.

注意第二种写法虽然发生在编译区域内,但被包装的 get_magic_num 本身仍然定义在 torch.compile 区域之外,这正符合第 3 节的要求 3。如果试图在编译函数内部定义 trace_me 再包装它,Dynamo 会抛出明确的 Unsupported 异常,错误信息为(见 test_decorators.py):

Applying `nonstrict_trace` to function <trace_me>; however, `nonstrict_trace`
currently requires the function to be defined outside `torch.compile` region.

4.3 更完整的实战形态:把整个子模块的前向包装起来

除了文档的最小示例,torch.compiler.nonstrict_trace 的 docstring 还给出了一个贴近实际训练场景的形态——把 nn.Module 作为输入传进去,并验证梯度可回传(见 torch/compiler/init.py):

import torch

@torch.compiler.nonstrict_trace
def traced_forward(model, x):
    # nonstrict_trace 区域内部允许出现 dynamo graph break
    torch._dynamo.graph_break()
    return model(x) + x

class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.inner = torch.nn.Linear(10, 10)

    def forward(self, x):
        return traced_forward(self.inner, x)

model = MyModule()
opt_model = torch.compile(model, backend="aot_eager", fullgraph=True)
out = opt_model(torch.randn(10, 10))
out.sum().backward()  # 梯度能穿过 traced_forward 流回 model.inner 的参数

这里 nn.Module 作为输入是被支持的:其参数与 buffer 会被纳入 autograd 追踪。

5. 源码解析:nonstrict_trace 是如何生效的

5.1 装饰器实现:包装函数 + 双名单注册

核心实现在 torch/_dynamo/decorators.py。它做了几件事:

def nonstrict_trace(traceable_fn: Callable[_P, _R]) -> Callable[_P, _R]:
    if not callable(traceable_fn):
        raise AssertionError("nonstrict_trace expects a callable")

    _check_mutually_exclusive_decorators(traceable_fn, "nonstrict_trace")

    @functools.wraps(traceable_fn)
    def wrapped(*args, **kwargs):
        return traceable_fn(*args, **kwargs)

    wrapped_id = id(wrapped)

    # 这行让它能复用 allow_in_graph 的大部分实现
    trace_rules._allowed_callable_ids.add(wrapped_id)

    # 这行让实现与 allow_in_graph 产生分叉
    trace_rules._nonstrict_trace_callable_ids.add(wrapped_id)

    # 避免 id 复用引发隐蔽 bug
    def deregister():
        trace_rules._allowed_callable_ids.remove(wrapped_id)
        trace_rules._nonstrict_trace_callable_ids.remove(wrapped_id)

    weakref.finalize(wrapped, deregister)
    return wrapped

可以读出三个设计要点:

  1. 包装而非原地修改:装饰器返回一个 functools.wraps 生成的新包装函数,并把该包装函数的 id 同时加入两份名单——_allowed_callable_ids(复用 allow_in_graph 的入图放行逻辑)与 _nonstrict_trace_callable_ids(标记为"非严格可追踪")。名单定义与查询函数位于 torch/_dynamo/trace_rules.py_nonstrict_trace_callable_ids()is_nonstrict_trace_callable(obj),后者仅判断 id(obj) in _nonstrict_trace_callable_ids)。
  2. weakref.finalize 防止 id 复用:注册表以 id() 为键,而 Python 中对象销毁后 id 可能被复用,因此装饰器为每个包装函数挂了一个 finalizer,在函数被垃圾回收时自动从两份名单中移除——这也是 allow_in_graph 采用的同款防御手段(decorators.py)。
  3. 与 leaf_function 互斥_check_mutually_exclusive_decorators 会检查函数是否同时被标记为 @leaf_function@nonstrict_trace,二者只能选其一(decorators.py)。

5.2 追踪时如何路由:AllowInGraphKind.NONSTRICT_TRACE

当 Dynamo 追踪到对已标记函数的调用时,走的是 TorchInGraphFunctionVariable 路径。对方法(method)形态的 nonstrict_trace 函数,torch/_dynamo/variables/functions.pyUserMethodVariable.call_function 有专门分支:

from ..trace_rules import is_leaf_function, is_nonstrict_trace_callable

if is_nonstrict_trace_callable(self.fn):
    call_args = [*self.self_args(), *args]
    var = variables.TorchInGraphFunctionVariable(
        self.fn, kind=variables.torch.AllowInGraphKind.NONSTRICT_TRACE
    )
    return var.call_function(tx, call_args, kwargs)

其中 self.self_args() 把方法的 self 也拼进参数列表。kind=AllowInGraphKind.NONSTRICT_TRACE 这个标记就是 Dynamo 与 allow_in_graph(默认 kind)行为的"分叉点":非严格追踪版本允许 pytree 注册的用户类作为输入/输出、允许 nn.Module 输入、把捕获张量按常量处理。

5.3 与 allow_in_graph 的关系

torch.compiler.allow_in_graph 的 docstring 明确写道:"See also nonstrict_trace(), which has slightly fewer restrictions on the inputs"(见 torch/compiler/init.py)。两者差异可归纳为:

维度 allow_in_graph nonstrict_trace
输入/输出类型 必须是 FX 图中 Proxy-able 的类型(Tensor/int/bool/float/None 及它们的 List/Tuple) pytree 兼容类型即可,用户自定义类需注册 torch.utils._pytree.register_pytree_node / register_dataclass / register_constant
nn.Module 输入 不支持 支持,参数/buffer 参与 autograd
捕获的全局/闭包张量 必须显式作为输入传入 可以捕获,但按常量处理、不建 guard、不回梯度
基本值与容器结构 每个调用点会被特化(specialized):同一调用点每次执行须保持相同的基本值与结构

选择建议(源自源码 docstring 对 leaf_function 的对比说明,同样适用于三者之间的取舍):若函数有静态计算图、无运行时副作用,优先 allow_in_graphnonstrict_trace,它们允许 AOT Autograd 追踪穿过并继续优化;若函数必须保留 eager 执行语义(日志、外部库调用等运行时副作用),则考虑 leaf_function

6. 行为边界与错误场景:从测试用例看

test/dynamo/test_decorators.py 中有一组 test_nonstrict_trace_* 用例,系统覆盖了该特性的行为边界,值得逐一对照理解:

支持的结构形态

  • 张量参数与嵌套调用test_nonstrict_trace_tensor_args 验证被追踪函数嵌入更大表达式(t0 * t2)时 aot_eager 结果与 eager 参考一致(L320-L337);test_nonstrict_trace_from_torch_compiler 进一步确认 torch.compiler.nonstrict_tracetorch._dynamo.nonstrict_trace 行为一致。
  • None 输入/输出test_nonstrict_trace_none_inputstest_nonstrict_trace_none_outputs 验证 None 可以出现在参数与返回值(tuple、dict 内)中(L358-L396)。
  • dict 与带副作用的容器test_nonstrict_trace_pre_existing_dicttest_nonstrict_trace_newly_constructed_dict_with_side_effects 验证既有 dict、新建 dict、以及调用前后对 dict 的修改都能得到与 eager 相同的结果(L398-L453)。
  • 用户自定义类:只要先通过 register_pytree_node 注册(测试基类 PytreeRegisteringTestCase 提供该方法),Point、嵌套的 PointTensor 等类实例都能作为输入输出;test_nonstrict_trace_nested_custom_class 还验证了 nonstrict_trace 函数内部再调用另一个带 graph_break() 的辅助函数也能被正确追踪(L456-L612)。
  • 方法形态test_nonstrict_trace_on_method 验证 @nonstrict_trace 可以直接装饰类方法,self 会自动作为隐式输入参与追踪(L729-L755)。
  • 捕获的外部张量按常量处理test_nonstrict_trace_captured_external_tensorcst = torch.ones(1) 是闭包捕获的全局张量,编译结果与 eager 一致,印证了第 3 节"非输入值视为常量"的语义(L757-L773)。
  • 符号值输出test_nonstrict_trace_tuple_and_sym_int_output 验证在 dynamic=True 下,函数返回 x.size(0)(SymInt)这类符号值也是可行的(L680-L695)。
  • nn.Module dict 输入test_nonstrict_trace_nn_module_dict_input 验证把 {"a": linear_a, "b": linear_b} 这样的模块字典传入被追踪函数后,modules"a" + modules"b" 能在 fullgraph=True 下正确编译运行(L1002-L1019)。

明确的"不生效"与"报错"边界

  • 作用域限定test_nonstrict_trace_no_action_at_a_distance 表明,对某函数调用 torch._dynamo.nonstrict_trace(trace_me)(拿到包装函数)却不使用返回值,对直接调用原函数没有任何效果——图中仍会出现 1 次 graph break(L775-L796)。注册是"按包装函数对象"生效的。
  • 区域内定义即报错test_nonstrict_trace_inside_compiled_function_error(上文第 4.2 节已引用),要求函数必须定义在编译区域外。
  • 未注册的用户类报错test_nonstrict_trace_custom_class_errortest_nonstrict_trace_nested_custom_class_error 中,Point 未注册 pytree 时报 Unsupported: Invalid input type for nonstrict_trace-ed functionL815-L886)。
  • 不支持的输出类型test_nonstrict_trace_custom_class_output_error 中,函数返回一个未注册的自定义类实例时报 Unsupported output type for nonstrict_trace-ed function——注意输入侧可以通过 pytree 注册解决,而输出侧该测试表明返回未注册类同样不被接受(L888-L914)。
  • 不透明对象的特殊处理test_nonstrict_trace_pre_existing_register_constant_type_guard 演示了与 torch._library.opaque_object.register_custom_class(State, typ="constant") 的交互:常量类对象作为输入时,相同实例不触发重编译(frame_count 保持 1),换值(State(42)State(41))才重编译(L614-L661)。

7. 选型速查:何时用 nonstrict_trace

综合文档与源码,可以形成如下判断路径:

  1. fullgraph=True 下某函数调用触发图断点,且该函数体由常规张量运算构成、没有必须 eager 的副作用 → 用 nonstrict_trace,让 AOT Autograd 把函数体展开进图;
  2. 函数输入/输出是普通张量与基础 Python 类型、且不想引入 pytree 注册 → 更简单的 allow_in_graph 即可;
  3. 输入/输出涉及用户自定义类或 nn.Module → 选择 nonstrict_trace,用户类需先注册 torch.utils._pytree
  4. 函数必须保留 eager 执行语义(日志、外部 I/O、数据依赖的运行时分支) → 不属于非严格追踪的适用范围,应考虑 custom op 或 leaf_function
  5. 函数读到的全局/闭包张量请按常量对待:不要指望对其回梯度,也不要在其后依赖函数对这些值的修改;
  6. 基本值(int/float 等)与容器结构在每个调用点上是特化的:同一调用点每次执行必须保持相同的基本值与结构,否则会触发重编译或错误。

8. 小结

torch._dynamo.nonstrict_trace 是 PyTorch 编译器编程模型中"让不可被 Dynamo 严格追踪的函数体仍进入最终编译图"的官方手段。它的机制是:装饰器将函数包装并在 trace_rules 的允许名单与 nonstrict 名单中按 id() 双注册(weakref 自动清理);Dynamo 追踪时以 AllowInGraphKind.NONSTRICT_TRACE 路由为不透明调用,把追踪责任交给 AOT Autograd;最终 FX 图包含函数内部全部相关张量操作,fullgraph=True 下不再有图断点,且梯度可以穿过该函数完成训练。使用时牢记四条前置要求——满足通用非严格追踪、输入输出为基本类型或已注册 pytree 类型、函数定义在编译区域外、捕获值按常量处理——即可在 test/dynamo/test_decorators.py 所覆盖的绝大多数场景下正确应用。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
927
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
603
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
397
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.04 K
525