PyTorch 动态形状(Dynamic Shapes)核心概念详解:SymInt、Guards 与编译期符号推理
2026/9/10 11:29:53 网站建设 项目流程

PyTorch 动态形状(Dynamic Shapes)核心概念详解:SymInt、Guards 与编译期符号推理

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

本文基于 PyTorch 官方文档 dynamic_shapes_core_concepts.md 整理并扩充,面向torch.compile编译器栈的工程者与进阶用户,系统讲解动态形状的两大核心原语——符号整数(SymInt)与守卫(Guards)——以及 Runtime Asserts、Hint 值、动态行为诊断与整体架构。读完后,你将能够理解一次torch.compile编译中符号形状是如何从"分配符号 → 算子传播 → 守卫约束 → 简化与安装"这条链路流动的,并能结合仓库源码定位mark_dynamicShapeEnvSymNode等关键实现。

1. 符号整数(SymInt):用代数表达"未定的尺寸"

符号整数(Symbolic Integers,简称 Symints)用于表示"可以跨取一个范围的变量",是动态形状体系的基础数据单元。文档给出的经典例子是:

x = torch.randn(5, 5) # shape: [5, 5] torch._dynamo.decorators.mark_dynamic(x, 0) x = torch.randn(5, 5) # shape: [s0, 5],第 0 维被符号化 y = torch.cat([x, x], dim=0) # shape: [2*s0, 5],形状以 Sympy 表达式呈现

可以看到,cat之后张量的形状不再是具体整数,而是一个符号表达式2*s0——形状计算在编译期以符号形式完成,而不是等到运行时才确定。

文档同时强调了一个容易被忽略的边界:z = x * y直接报错,因为我们知道逐点乘法要求两个张量形状相同,而编译期可以静态证明s0 != 2 * s0。细心的读者会指出:当s0 == 0时两者其实相等。文档给出的解释是——这种特例之所以被安全忽略,源于 PyTorch 的zero-one specialization(0/1 特化)机制,详见姊妹篇 动态形状的 0/1 特化。

1.1 mark_dynamic 的源码行为

mark_dynamic定义在 torch/_dynamo/decorators.py,签名为:

def mark_dynamic( t: Any, index: int | list[Any] | tuple[Any], *, hint_override: int | None = None, min: int | None = None, max: int | None = None, specialize_on: list[Any] | None = None, ) -> None:

从源码注释和实现可以确认文档中"如果已知min/max可以指定"的说法:

  • min/max会被写入张量的_dynamo_unbacked_bounds属性,直接收窄该符号的取值范围(见 decorators.py 中 bounds 的写入逻辑);
  • index可以是单个维度、维度列表或元组,支持一次标记多个维;
  • 该 API 被标注@forbid_in_graph——即所有mark_dynamic调用必须在torch.compile之前完成,尝试在可 trace 的函数内调用会显式抛错;
  • 行为受torch._dynamo.config.dynamic_shapes控制:配置为False时与mark_dynamic并用会抛异常(docstring 注明该支持"将最终实现");
  • 进阶参数hint_override会替换首次示例输入推导出的 size hint,且会影响 Inductor 代码生成决策(autotuning、归约策略),因此会进入FxGraphCache的缓存键(通过ShapeEnv.var_to_hint_override);
  • 进阶参数specialize_on支持"一次泛型 trace + 多份特化编译",例如specialize_on=[lambda x: x == 8, lambda x: x == 16]会生成一份 Dynamo trace 和两次后端编译,运行时输入命中条件即派发到特化版本。

此外,mark_dynamic与"弱标记"maybe_mark_dynamic存在优先级规则:对同一维度,后者的标记会被mark_dynamic接管并强制执行更强的动态语义。

2. Guards:保证已编译代码图仍然有效

torch.compile中,guard(守卫)是"确保编译代码图依然有效"的机制。把某个变量动态化后,其默认取值范围是[-inf, inf],此时任何输入尺寸都复用同一份编译产物;但一旦函数内部出现了依赖具体数值的分支,就必须把分支条件记录成 guard。文档中的例子:

def foo(x): if x > 5: return x / 2 return x / 3

第一次调用foo(6)x / 2分支,编译器随之添加 guardx > 5;之后再调用foo(4)时该 guard 求值为假,触发重新编译。这就是"守卫破坏 → recompile"的基本模型。

2.1 守卫在架构中的位置

文档"Overall Architecture"一节给出的六步流程中,guard 出现在第 4、5、6 步:

  1. Dynamo 编译一个 frame 时,分配一个ShapeEnv(挂在FakeTensorMode上)来跟踪符号形状;
  2. 根据策略决定在入口处为张量分配符号尺寸;
  3. 符号尺寸随算子传播,同时维护两套表示:用于符号计算导出的FX IR与用于推理的Sympy 表达式
  4. Dynamo trace 期间或 Inductor 优化期间产生的条件会转成 guard,来源既有 Python 也有 C++;
  5. guard 可以反过来简化符号变量——例如断言s0 == 4之后,所有s0的出现都可以被4替换;
  6. trace 与优化完成后,所有 guard 随编译产物一起安装,只有当全部 guard 求值为真时才允许复用。

值得注意的是第 5 步:guard 不只是"失效检测器",它同时是符号简化的信息源,这解释了为什么守卫系统必须保持精确——一条错误安装的守卫会让后续推理建立在错误的前提上。guard 的实际求值与安装代码集中在 torch/_dynamo/guards.py,而符号层面的守卫记录则落在ShapeEnv(见 torch/fx/experimental/symbolic_shapes.py)。

3. Runtime Asserts:向编译器提供已知事实

当你明确知道某些事实(例如 batch size 一定小于 100)时,可以用 runtime assert 主动告知编译器:

def foo(batch_size): torch._check(batch_size < 100) if batch_size < 100: return do_something return do_something_else()

torch._check引入的断言会让系统做三件事(文档"Value Ranges and Constraints"一节的描述):

  1. 尝试用等价表达式替换未回溯(unbacked)符号;
  2. 依据断言收窄符号的取值范围(value range refinement);
  3. 记住那些恒为真的布尔表达式,供后续守卫与简化使用。

这与 guard 的区别在于:guard 是从条件分支被动"诱发"出来的约束,而 runtime assert 是用户主动注入的先验知识,能在没有 Python 分支的情况下同样驱动符号推理。

4. Hint 值:编译期的"具体样例"

"Hint value" 指编译过程中实际已知的具体数值。符号形状虽然以 Sympy 表达式参与计算,但 JIT 编译器在做表达式决策时(例如选择归约策略、判断循环上界的可行性)常常需要一个具体数字——hint 就提供了这个"编译时样例值",使得不同维度的调用可以复用同一份编译产物而无需反复重编译。

这与mark_dynamichint_override参数直接对应:不传时 hint 默认取首次示例输入在该维度上的实际值;显式传入hint_override则替换掉这个默认 hint,并因此进入 FxGraphCache 的缓存键(因为 hint 会影响 Inductor 的 autotuning 与归约策略选择)。

5. 动态行为的整体工作方式(Dynamic Behavior Overview)

文档对"动态形状到底何时发生、如何诊断、如何控制"给出了完整的操作层面总结,这里逐条展开:

5.1 默认静态、尺寸变化才动态

PyTorch默认假设静态形状。当检测到尺寸变化时,Dynamo 会尝试以动态输入重新编译,但若存在条件分支或缺少对动态形状的支持,这次重编译可能失败。要诊断"过度特化"(overspecialization),可以设置TORCH_LOGS=dynamic,观察日志中的 "eval" 条目——它们指示守卫是何时、因为什么被添加的。日志格式本身的说明见 动态形状调试:tlparse 与 TORCH_LOGS。

5.2 提前标记 vs 两种 dynamic 开关

  • 预期某维会是动态时,用torch._dynamo.mark_dynamic(tensor, dim)提前标记,已知上下限时同时给出min/max
  • torch.compile(dynamic=False)关闭自动动态形状,每个新尺寸都会触发一次重编译——简单、可预测,但编译次数随尺寸增长;
  • torch.compile(dynamic=True)尽可能多地使用动态形状,最适合小型模型;文档明确提醒它"对大型模型未必合适,可能带来崩溃或性能问题"。

5.3 按来源白名单:dynamic_sources 与 static_sources

对含图间断点(graph breaks)的大模型,有时很难找到"该动态标记哪些输入"。此时可以按来源(source)名做白名单,且由于 source 名在图间断前后保持稳定,动态性可以跨断点保持。

文档提到的两个变量与配置对(在 torch/compiler/config.py 中定义):

  • 动态白名单:环境变量TORCH_COMPILE_DYNAMIC_SOURCEStorch.compiler.config.dynamic_sources。取值为逗号分隔的 source 名列表,例如"L['x'], L['y']",也支持正则,例如"L\['x.*'\], L\['y.*'\]";它甚至可以把普通整数标记为动态。这个白名单优先级高于dynamic=Falseforce_nn_module_property_static_shapesforce_parameter_static_shapes等其它开关。
  • 静态镜像TORCH_COMPILE_STATIC_SOURCES/torch.compiler.config.static_sources把列出的来源钉在静态上,接受同样的 source 名、正则与:N逐维语法,且优先于自动动态形状、PGO 与dynamic=True。当 PGO 或dynamic=True误判了某来源、而你希望它保持静态时,这是对应的"反向阀门"。

5.4 eager_then_compile stance:让框架替你推导动态性

文档指出的另一个务实选项是eager_then_compilestance——愿意接受首个 batch 的性能代价,换取框架自动推导出哪些输入该动态。入口是torch.compiler.set_stance(实现见 torch/compiler/init.py)。从 docstring 可确认可选 stance 包括:

stance行为
default正常编译
force_eager忽略所有torch.compile指令
eager_on_recompile需要重编译时退化为 eager,命中缓存的编译产物仍照常使用
fail_on_recompile触发重编译即抛错
eager_then_compile首次调用 eager、后续调用编译;从前两次调用的差异中推断动态性,避免第一次调用浪费在静态编译上
aot_eager_then_compile首次调用走 AOT eager(可获得 activation checkpointing 的内存收益),后续编译

set_stance可作函数、上下文管理器或装饰器使用,但不能在torch.compile区域内调用。

6. 整体架构:符号形状的五步工作流

文档"Overall Architecture"一节给出了符号形状的端到端工作流,这是理解整套机制的骨架:

  1. ShapeEnv 分配:Dynamo 编译一个 frame 时分配ShapeEnv(挂在FakeTensorMode上),负责跟踪符号形状;
  2. 入口符号化:根据策略决定在入口处为张量分配哪些符号尺寸;
  3. 符号传播:符号尺寸穿过算子传播,同时维护 FX IR(供符号计算导出)与 Sympy 表达式(供推理)两套表示;
  4. 守卫归纳:Dynamo trace 或 Inductor 优化期间产生的条件转成 guard,来源横跨 Python 与 C++;
  5. 守卫简化符号:断言s0 == 4之后,所有s0都可以替换为4
  6. 守卫安装:trace 与优化结束后,全部 guard 随编译代码安装,全部为真才允许复用。

配套文档可以从两个方向继续深入:动态形状入门(backed/unbacked 的区别)与动态形状故障排查。

7. 内部 API 类层次

文档给出了 Python 与 C++ 两侧对照的类层次,这是阅读动态形状源码的导航图:

7.1 Python 侧

  • SymInt/SymFloat/SymBool:用户可见类,模拟int/float/bool的行为。两个SymInt相加产生一个新的SymInt,符号化地跟踪这次整数加法;
  • SymNode:内部结构(可通过symint.node访问),保存实际的符号跟踪信息。SymNode类型擦除的,因此方便表达混合类型的运算;
  • ShapeEnv:每次编译一份的上下文状态,跟踪所有自由符号与迄今累积的全部守卫。每个SymNode都记录它所属的ShapeEnv,但反向不成立——SymNode只有在参与某个 guard 时才会被保留使用。

7.2 C++ 侧

  • c10::SymInt/SymFloat/SymBool:与 Python 对应,模拟int/float/bool
  • c10::SymNode/SymNodeImpl:对应 Python 的SymNode
  • 没有 C++ 版 ShapeEnv:为便于调试,整套符号推理设施留在 Python 一侧。

由此可推断一条明确的开发约束:任何希望被make_fx等工具 trace 的代码,都必须能处理流经它的SymInt/SymFloat/SymBool——例如不能对符号尺寸做int(x)之外的假设性断言,也不能对 Sympy 表达式使用只对普通整数成立的分支。

8. 取值范围与约束(Value Ranges and Constraints)

符号变量维护取值范围,描述其可能的取值集合。文档给出的默认值是:

  • 尺寸类的 unbackedSymInt:取值范围[0, Inf]
  • 普通 unbackedSymInt:取值范围[-Inf, Inf]

当断言发生(如torch._check(x == y))时,系统依次执行:

  1. 尝试用等价表达式替换 unbacked 符号;
  2. 依据断言收窄取值范围;
  3. 记住恒为真的布尔表达式。

min/max参数(mark_dynamic)、torch._check断言、guard 三者都是对取值范围这一中心数据结构的写入路径,这也是第 5 步"guard 简化符号"能够成立的基础:范围收窄到单点即可执行常数替换。

8.1 关键文件索引

文档末尾给出的"Important files"清单,结合仓库实际路径整理如下:

关注点文件
C++ SymInt APIc10/core/SymInt.h,同目录SymFloat.hSymBool.h
Python SymInt APItorch/init.py(查找SymInt/SymFloat/SymBool
C++ 衔接层(plumbing)c10/core/SymNodeImpl.h、torch/csrc/utils/python_symnode.h、torch/csrc/jit/python/init.cpp
Python 基础设施(ShapeEnv/SymNode 核心)torch/fx/experimental/symbolic_shapes.py(ShapeEnv定义于 该文件第 3957 行 附近)
其它重要文件torch/_subclasses/fake_tensor.py、torch/_meta_registrations.py、各算子的 decomps 与 PrimTorch refs

9. 小结:从文档脉络到源码入口

把整份核心概念文档串起来,动态形状的机制可以浓缩为一句:ShapeEnv在每次编译中建立符号世界,算子传播产生 Sympy 表达式,Python/C++ 两侧的条件沉淀为 guards,guards 再反过来收窄符号范围、简化表达式,最终 guard 集合作为复用条件随编译产物安装。用户侧的三个杠杆——mark_dynamic(含min/max)、torch._checkruntime assert、dynamic_sources/static_sources白名单——本质上都是在向这套符号推理系统注入先验。

进一步阅读建议按此路径展开:先读 动态形状基础(backed/unbacked) 与 进阶用法 巩固 API 层;遇到"守卫破坏/重编译"问题时查 Troubleshooting 与 GuardON 错误;想理解 0/1 特化为何让s0 == 0的边界情况无关紧要,读 0/1 特化专章;动手排查时用TORCH_LOGS=dynamic配合 tlparse/TORCH_LOGS 调试指南 查看 "eval" 条目;最后以 symbolic_shapes.py 与 torch/_dynamo/guards.py 作为源码级入口精读实现。

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询