GraphMend: Code Transformations for Fixing Graph Breaks in PyTorch 2
A compiler technique that automatically fixes FX graph breaks in PyTorch 2 programs through AST-level program analysis and transformations
GraphMend: 通过代码变换自动修复 PyTorch 2 中的图断裂
一、论文概述
| 项目 | 内容 |
|---|---|
| 标题 | GraphMend: Code Transformations for Fixing Graph Breaks in PyTorch 2 |
| 作者 | Savini Kashmira, Jayanaka Dantanarayana, Thamirawaran Sathiyalogeswaran, Krisztian Flautner, Lingjia Tang, Jason Mars |
| 机构 | 密歇根大学 (University of Michigan) |
| 论文 | https://arxiv.org/abs/2509.16248 |
| 发布 | 2025-09-17 (v1), 最新 v4: 2026-06-29 |
| 许可 | CC BY 4.0 |
| 领域 | 编程语言 (cs.PL); 机器学习 (cs.LG); 软件工程 (cs.SE) |
二、核心思想
问题定义
PyTorch 2 引入了 TorchDynamo 和 TorchInductor 来实现即时(JIT)图编译,但某些代码模式仍会导致 FX 图断裂(graph breaks),迫使执行回退到 Python 急切模式。这引入了昂贵的 CPU-GPU 同步开销并减少了优化机会。
我们对 195 个 Hugging Face 模型的调查显示:13.8% 的模型存在图断裂,共计 147 个断裂实例。其中包括大量广泛使用的模型:Google 的 T5-small(月下载量超 280 万)、OpenAI 的 Whisper 系列、Meta 的 BART 系列等。
解决方案概述
GraphMend 是一种编译器技术,通过源代码级别的分析与变换自动消除可修复的图断裂。它在 Python 字节码生成之前,在 AST 层面对程序进行分析和变换,使 PyTorch 能够捕获更大、不间断的 FX 图,无需开发者手动重构代码。
GraphMend 在 Jaseci 框架内实现——一个接受标准 Python 程序并将其编译为原生 Python 字节码的前端编译器。GraphMend 扩展了标准 PyTorch 2 编译管道,增加了一个 AST 级变换阶段。
核心成果:在全部 27 个存在图断裂的模型上评估,GraphMend 消除了 147 个断裂中的 107 个(73%),在 21 个模型中完全修复了所有断裂。在 NVIDIA GPU 上实现高达 26× 冷启动加速(平均 5×),稳态前向传播最高加速 1.39×。
三、技术架构
整体框架

GraphMend 集成在 Jaseci 编译管道中,工作流程如下:
Python 源代码
│
▼
Jaseci 解析 → AST → CFG + 符号表 → UniiR (统一中间表示)
│
▼
GraphMend 编译器分析:
┌─────────────────────────────────────┐
│ 1. Dynamo 入口点与断裂检测阶段 │
│ - 定位 torch.compile 装饰的函数 │
│ - 标记图断裂候选及其类别 │
│ 2. AST 变换阶段 │
│ - 应用三种变换规则消除断裂 │
│ - 语义等价性验证 │
│ 3. 生成变换后的字节码 │
└─────────────────────────────────────┘
│
▼
标准 Python 字节码 → PyTorch 2 TorchDynamo → FX 图 → TorchInductor
核心公式:三种变换规则
GraphMend 定义了三种守恒语义的变换规则,每种都有严格的合法性条件:
规则 1: Predicated Trap Lowering ([Trap])
将 Python 条件异常转换为图内张量断言:
if not C: raise E(msg)
⇒
torch._assert_async(tensorize(C), uid + msg)
合法性条件:guard 主体仅包含 raise 语句,且 tensorize(C) 有定义。
GraphMend 通过唯一标识符前缀生成的断言消息,保留原始异常类型和消息。因为 torch._assert_async 抛出 RuntimeError,处理器识别标识符并重新抛出原始异常类型和消息。
规则 2: Predicated Data-dependent Control Flow ([Where])
将数据依赖的条件分支转换为 torch.where:
if c: St; K[et]
else: Sf; K[ef]
⇒
p = c; hoist(St); hoist(Sf);
K[torch.where(p, et, ef)]
合法性条件:
- 分支值 et、ef 是具有匹配形状和 dtype 的张量表达式,且都不为 None
- 区域内每个函数调用通过 UniiR 解析为可检查的代码
- 分析验证其无可见副作用、异常和图断裂诱导器
Setup 语句可提升的条件:不包含 return/break/continue;未提升路径上不会引发;任何写入必须不可观察(写入值仅在分支连接处读取,或在未对侧路径上通过逃逸检查验证,或是幂等操作如设备移动)。
规则 3: Deferred Side Effect ([Defer])
将副作用操作延迟到 FX 图执行后:
S1; f(args); S2; return e
⇒
S1; buf.push(f, clone(args)); S2;
r = e; flush(buf); return r
合法性条件:被调用方解析为内置 print 函数或受支持的 logger,且未被重新绑定。其唯一可见效应是控制台或日志输出。
在原始调用点,GraphMend 记录被调用方及其已求值的参数。张量参数被克隆以抵御后续变异。记录的调用按 FIFO 顺序在编译张量计算完成后重放。
变换顺序
[Trap] → [Where] → [Defer]
- [Trap] 优先:因为分支局部的 raise 会阻止控制流转换
- [Where] 其次:因为它可能重组分支、合并调用并引入新的返回路径
- [Defer] 最后:在最终的控制和返回结构已知后应用
语义等价性证明
GraphMend 保证可观测行为不变:张量输出、可见状态、控制台和日志输出的内容与顺序、以及抛出的异常及其类型和消息。每个变换在不改变任何可观测内容的前提下仅改变内部实现:
- [Where] 仅添加不引发异常且写入无副作用的工作
- [Trap] 仅改变异常报告的时序
- [Defer] 仅延迟不改变最终结果的输出
四、核心创新
| 创新点 | 说明 | 理论/实验依据 |
|---|---|---|
| AST 级图断裂检测 | 在字节码级别之上分析源代码结构,识别 5 类断裂模式(数据依赖控制流、日志/打印、验证守卫、动态形状、tensor.item()) | 从 195 个 HF 模型中识别出 27 个存在断裂的模型 |
| 三种守恒变换规则 | [Trap]/[Where]/[Defer] 在严格合法性条件下消除断裂,保证语义等价 | 消除 107/147 (73%) 断裂 |
| UniiR 统一中间表示 | 将 AST、CFG 和符号表合并为统一 IR,支持跨层次的程序分析 | 使函数调用解析和副作用分析成为可能 |
| Jaseci 框架集成 | 在标准 Python 字节码生成前插入变换阶段,无需修改用户代码 | 用户只需安装 PyPI 包即可运行 |
五、代码实现分析
GraphMend 基于 Jaseci 框架实现:
- 入口:作为 PyPI 包分发,用户安装后可直接运行模型
- 分析阶段:
Dynamo Entry-Point and Break Analysis Pass- 定位被
torch.compile包装或装饰的函数/nn.Module 对象 - 标记每个图断裂候选及其类别
- 谓词分类:数据依赖(truth value 取决于运行时张量内容)vs 静态(可从形状/常量推断)
- 定位被
- 变换阶段:
AST Transformation Pass- 按 [Trap] → [Where] → [Defer] 顺序应用变换
- 每个变换通过 UniiR 上的合法性检查
- 输出:生成无断裂的字节码,传递给标准 PyTorch 2 编译管道
六、实验结果
数据集
从 Hugging Face 选取 195 个模型:前 100 个热门模型 + 前 30 个下载量最高模型 + 65 个随机采样模型(覆盖 diverse architectures 和 task categories)。其中 27 个 (13.8%) 存在图断裂。
图断裂修复能力
| 修复情况 | 模型数 | 说明 |
|---|---|---|
| 完全修复 | 21/27 | 所有可修复断裂均已消除 |
| 部分修复 | 3/27 | 部分断裂被修复 |
| 未修复 | 3/27 | 断裂源于 fundamentally untraceable 操作(如动态形状算子) |
| 总计消除断裂 | 107/147 (73%) |
断裂原因分类
| 类别 | 缩写 | 说明 |
|---|---|---|
| Data-dependent Control Flow | DC | 数据依赖的控制流 |
| Logger/Print Calls | LC | 日志/打印调用 |
| Validation Guards | VG | 验证守卫(如 assert) |
| Dynamic Shape Operators | DS | 动态形状算子 |
| Tensor Item Extraction | TI | tensor.item() 调用 |
冷启动加速(Cold Start Speedup)
冷启动加速在 serverless 部署(如 AWS Lambda)和自动扩缩容场景下至关重要。
| 模型 | 断裂数 | 断裂原因 | 修复率 | RTX 3090 | A40 | H100 |
|---|---|---|---|---|---|---|
| bart-large-cnn | 7 | DC(4) + LC(3) | 100% | 21.07× | 20.22× | 24.1× |
| MoLFormer-XL-both-10pct | 5 | VG(5) | 100% | 24.71× | 24.25× | 25.92× |
| rebel-large | 7 | DC(4) + LC(3) | 100% | 19.86× | 21.90× | 23.20× |
| Florence-2-large | 7 | DC(7) | 100% | 20.95× | 19.55× | 4.39× |
| grounding-dino-tiny | 17 | DS(3)+DO(3)+DC(11) | 58% | 5.20× | 5.19× | 5.50× |
| 平均 | - | - | 73% | 5× | 5× | 5× |
关键发现:数据依赖控制流断裂的模型获得最大收益。因为每个此类断裂强制设备同步和 D2H memcpy,使 GPU 空闲。修复后前向传播作为单一连续 CUDA 图执行。
稳态前向传播加速(Steady State)
| 模型 | RTX 3090 | A40 | H100 |
|---|---|---|---|
| tiny-random-PegasusForCausalLM | 1.34× | 1.37× | 1.39× |
| Florence-2-large | 1.19× | 1.21× | 1.23× |
| bart-large-cnn | 1.13× | 1.11× | 1.17× |
| 平均 | ~1.1× | ~1.1× | ~1.1× |
稳态加速较温和,因为 CUDA Graph 在首次执行后已被录制,后续执行通过重放加速。消除图断裂的主要收益来自跨断裂边界的算子融合。
吞吐量提升
| 模型 | 吞吐量增益 | 说明 |
|---|---|---|
| Florence-2-large (VLM) | 15% | 非自回归模型,前向传播占总推理时间比例大 |
| Qwen-Audio-Chat | 8% | 自回归模型,解码循环受内存带宽限制 |
| Phi-4-mini-instruct | 6% | 自回归模型 |
吞吐量增益小于稳态前向传播增益,因为端到端推理涉及更多环节(tokenization、采样、CPU 调度)。根据 Amdahl 定律,改进受限于前向传播在总时间中的占比。
图断裂性能影响分析

原始模型中,每次图断裂导致:
cudaStreamSynchronize— CPU 阻塞直到第一个图完成- Device-to-Host memcpy — 传输张量值用于 Python 条件判断
- CPU 端 eager mode 评估 if 条件
期间 GPU 完全空闲。修复后,整个前向传播作为单个 CUDA 图执行,GPU 持续满载。
冷启动时的额外开销:原始模型中每个 CUDA 图需要独立的录制阶段,GPU 空闲时间远超活跃时间。修复后仅单个录制阶段。

编译开销

GraphMend 自身的编译开销相对于标准 Python 执行极小,在 24 个测试模型中几乎可以忽略。
完整模型图捕获
对于 vLLM 和 SGLang 等推理框架,要求 torch.compile(fullgraph=True) 捕获完整图。21 个完全修复的模型均可满足此要求,使它们能够使用这些框架的编译和 CUDA Graph 重放路径。
七、相关工作
| 方面 | GraphMend | TorchDynamo-only | TorchScript |
|---|---|---|---|
| 分析级别 | AST + CFG + 符号表 | 字节码 | AST + 静态分析 |
| 变换能力 | 自动源代码变换 | 无 | 有限 |
| 语义验证 | 保守合法性检查 | 守卫函数 | 全部或无 |
| 用户改动 | 无需修改代码 | 可能需要 | 需要重写 |
八、总结
核心贡献
- 大规模实证研究:对 195 个 Hugging Face 模型的图断裂普及率调查,量化了生产中模型的问题范围
- 源代码级分析:识别可修复的图断裂模式,利用 AST、CFG 和符号表建立保守的合法性条件
- GraphMend 编译器技术:三种自动化 AST 级变换消除图断裂,在字节码生成前应用
- 语义等价保证:通过严格的合法性检查证明每个变换保持可观测行为不变
- 显著性能提升:26× 冷启动加速(平均 5×),1.39× 稳态加速,21/27 模型完全修复
局限性
- tensor.item() 调用:从张量提取 Python 标量,结果流入 Python 解释器,无论怎样表达都无法留在 FX 图内
- 动态形状算子:如
torch.nonzero或布尔掩码索引,输出形状数据依赖,无法通过源代码变换解决 - 3/27 模型完全未修复:断裂源于上述根本上无法追溯的操作
技术影响
GraphMend 证明了语义感知的源代码级分析和变换是 PyTorch 动态 JIT 编译管道的有效补充,大幅提升了可用性和性能。它使模型开发者无需手动重构代码即可获得显著加速。
九、参考资源
- 论文: https://arxiv.org/abs/2509.16248
- Jaseci 框架: https://github.com/jaseci-labs/jaseci
- PyTorch 2: https://pytorch.org/get-started/pytorch-2/
- torch._dynamo.config.reorderable_logging_functions: https://pytorch.org/docs/stable/dynamo.html
- vLLM: https://github.com/vllm-project/vllm
- SGLang: https://github.com/sgl-project/sglang