PyTorch 编译器 IR 指南:Core Aten IR 与 Prims IR 详解

📅 发布时间:2026/9/10 10:14:22
PyTorch 编译器 IR 指南:Core Aten IR 与 Prims IR 详解
PyTorch 编译器 IR 指南Core Aten IR 与 Prims IR 详解【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch导读PyTorch 2.0 引入了基于torch.compile的编译器技术栈为了统一后端对接方式PyTorch 为编译器后端提供了两套中间表示IRCore Aten IR与Prims IR。本文以 torch.compiler_ir.md 为骨架结合本仓库中torch/_prims、torch/_refs、torch/_decomp与torch/_inductor等模块的源码系统讲解这两套 IR 的设计定位、算子集合、底层实现原理以及它们如何在 Inductor 等后端中被消费帮助读者理解 PyTorch 编译器前端的核心抽象。一、两套 IR 的总体定位PyTorch 2.0 为后端提供了两套可对接的 IRCore Aten IRaten核心算子的子集保持函数式functional语义贴近现有 aten 算子体系主要用于与后端对接Prims IR更低层级的原语primitive算子集将高层算子进一步分解为显式类型提升与显式广播专门为编译器后端设计。两者均处于活跃开发中文档明确标注 This opset is still under active development, more ops will be added in the future算子数量会随版本持续扩充。二、Core Aten IR面向后端的函数式算子集2.1 设计目标Core aten ops 是aten算子中可以用来组合出其他算子的核心子集具有以下关键特征全函数式fully functional该算子集中不存在inplace变体如add_和_out变体如add_out所有算子都返回新结果语义干净便于编译器做图优化与重排复用现有算子定义与 Prims IR 不同Core Aten IR直接复用 native_functions.yaml 中已有的 aten 算子不另起炉灶不进一步分解Core Aten IR 不会把算子进一步拆解成显式的类型提升type promotion与广播broadcasting算子类型提升与广播逻辑仍然隐含在高层算子内部服务对象文档明确指出该算子集被设计为与后端对接的函数式 IR。2.2 从源码看 Core Aten 分解机制Core Aten IR 的算子集合并非手工维护的静态列表而是由分解decomposition机制动态确定。核心实现位于 torch/_decomp/init.pydef core_aten_decompositions() - CustomDecompTable: from torch.export.exported_program import default_decompositions return default_decompositions()core_aten_decompositions()返回导出流程中使用的默认分解表。而_core_aten_decompositions_post_autograd()函数同文件第 311 行起注释为NOTE [Core ATen Ops]给出了更精细的约束该列表复制自 torch/_inductor/decomposition.py并排除了会产生 prim 算子的分解最终分解得到的算子集合即 Core Aten 算子集。同时torch/_inductor/decomposition.py 中decompositions {**core_aten_decompositions(), **inductor_decompositions}表明 Inductor 的分解表由 Core Aten 分解表与 Inductor 自定义分解共同构成这印证了文档所述Core Aten IR 作为后端对接的公共函数式接口这一设计。2.3 算子清单的呈现方式文档中以 CSV 表形式列出了两套 IR 的算子清单引用路径为../../../build/ir/aten_ops.csvCore Aten IR 算子表../../../build/ir/prims_ops.csvPrims IR 算子表这两个 CSV 文件由构建过程build 阶段根据当时的分解配置生成因此在当前源码目录中并不直接存在。读者若想获取当前版本的完整算子清单需要在完成源码构建后查看build/ir/目录下的这两个文件该目录内容会随算子集演进而更新。三、Prims IR显式类型提升与广播的原语层3.1 设计目标Prims IR 是比 Core Aten IR更低一层的算子集同样可用于组合出其他算子其核心区别在于进一步分解Prims IR 会把高层算子进一步分解为显式的类型提升算子prims.convert_element_type与显式的广播算子prims.broadcast_in_dim面向编译器后端该算子集被设计为与编译器后端如 Inductor对接分解后的图更规整、更利于代码生成与向量化。3.2 两个核心原语的定义prims.broadcast_in_dim与prims.convert_element_type均在 torch/_prims/init.py 中通过_make_prim工厂函数定义。broadcast_in_dim第 1391 行附近的 schema 为broadcast_in_dim(Tensor(a) a, SymInt[] shape, int[] broadcast_dimensions) - Tensor(a)其语义是把输入张量a广播到目标shapebroadcast_dimensions指定原张量各维映射到新形状中的哪些维度。对应的_broadcast_in_dim_meta第 1280 行起实现了一套严格校验规则a.ndim必须等于len(broadcast_dimensions)每个维度都必须被映射len(shape)必须不小于a.ndim广播后的形状维度数不少于原张量broadcast_dimensions必须是严格升序的整数序列不允许维度相对重排每个映射维度必须落在新形状范围内。这套约束保证了广播操作不会改变维度顺序只做补维 拉伸是编译器可以安全进行布局与内存规划的前提。convert_element_type第 2036 行附近的 schema 为convert_element_type(Tensor a, ScalarType dtype) - Tensor它把张量的元素类型显式转换为目标dtype将 PyTorch 隐式类型提升如 int8 float32 自动提升为 float32在图中显式化。3.3 显式分解的价值在高层算子如add、mul的参考实现中类型提升与广播逻辑被炸开成上述两个原语例如torch/_refs/__init__.py中的实现大量调用broadcast_in_dim见第 1604 行附近的broadcast_in_dim(a, new_shape, broadcast_dimensions)模式。显式化带来两个直接好处后端无需重复实现类型提升/广播规则编译器后端只要实现两个原语的 lowering即可获得所有组合算子的支持图结构更规整每个算子输入输出的 shape/dtype 在图中一目了然便于做代数化简、公共子表达式消除与向量化。四、Inductor 如何消费 Prims IRInductor 是 PyTorch 的默认编译后端其 lowering 层直接为 Prims 原语注册了代码生成规则torch/_inductor/lowering.py 第 1086 行register_lowering(prims.convert_element_type, type_promotion_kindNone)对prims.convert_element_type注册了专门的 lowering 实现同文件第 1485 行register_lowering(prims.broadcast_in_dim, type_promotion_kindNone)对prims.broadcast_in_dim注册 lowering第 2869、2891、2904 行附近还有针对prims.convert_element_type的专项优化与模式匹配逻辑。从源码结构看Inductor 把broadcast_in_dim的实现建立在expand等底层视图操作之上把convert_element_type映射为元素级类型转换内核。这正印证了文档中Prims IR 被设计为与编译器后端对接的定位——两个核心原语是后端接入的最小契约。五、Core Aten IR 与 Prims IR 的选择建议维度Core Aten IRPrims IR层级较高贴近 aten 算子较低贴近原语算子来源复用native_functions.yaml中的 aten 算子独立的torch._prims原语集合类型提升隐含在高层算子内显式为convert_element_type广播隐含在高层算子内显式为broadcast_in_dim变体无inplace/_out变体全函数式全函数式原语典型消费方导出/对接后端时的通用函数式图Inductor 等编译器后端的 lowering 层实际使用中两种 IR 并非互斥导出与分解流程torch/_decomp/init.py先得到 Core Aten 级别的函数式图需要进一步下沉时再由分解表torch/_inductor/decomposition.py把算子分解到 Prims 层最终由 Inductor 的 lowering 层逐个原语生成代码。理解这一分层是深入阅读torch.compile全链路FX 图 → 分解 → lowering → 代码生成的关键起点。六、小结Core Aten IR是复用现有 aten 算子、全函数式、不含显式类型提升/广播的对接用 IR算子集合由core_aten_decompositions()分解表驱动并随构建产物输出在build/ir/aten_ops.csvPrims IR是更低层的原语集以broadcast_in_dim与convert_element_type为核心原语将类型提升与广播显式化被 Inductor 等后端直接 lowering两套 IR 均处于活跃演进阶段算子集合会持续扩充具体清单以构建生成的build/ir/目录下 CSV 文件为准。后续若需对接自定义编译器后端推荐以 Core Aten IR 作为函数式接口起点再依据后端能力选择是否下沉到 Prims IR如需深入了解可继续阅读 torch.compiler_ir.md 与 torch/_prims/init.py。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考