夜雨聆风学习资料网

ARTICLE · 1086635

LLM软件框架中的算子转化机制:从计算图到硬件执行的技术剖析

LLM软件框架中的算子转化机制:从计算图到硬件执行的技术剖析

摘要

大语言模型(LLM)的软件栈需要将高层次的神经网络描述(如 Transformer 结构)逐层转化为可在 GPU/NPU 上执行的算子(operator/kernel)。这一过程涉及前端图捕获、中间表示(IR)分层、算子分解与融合、后端代码生成等多个阶段。本文系统梳理主流 LLM 框架(PyTorch 2.x、JAX/XLA、TVM、Triton 生态)如何将 AI workload 转化为不同粒度的算子,并结合 GPU/NPU 底层执行模型分析这一转化链条的工程实现。


1. 引言

一个 LLM 训练/推理任务,从用户角度看只是 model(input) 这样一行 Python 代码;但从执行角度看,它需要被展开为成百上千个具体的计算核(kernel),每个 kernel 对应特定的硬件指令序列(如 Tensor Core 的 MMA 指令、ROCm 的 MFMA 指令)。中间的转化链条大致可以概括为:

Python/PyTorch 描述
   → 计算图捕获(Graph Capture)
   → 高层算子表示(ATen/HLO)
   → 算子分解(Decomposition)
   → 中间表示降级(IR Lowering,如 Linalg/TOSA)
   → 调度与融合(Scheduling & Fusion)
   → 后端代码生成(Codegen:Triton/LLVM/PTX/HSACO)
   → 硬件执行

下面按此顺序展开。


2. 图捕获层:从动态图到计算图

2.1 PyTorch 的路径:TorchDynamo + AOTAutograd

PyTorch 2.x 通过 torch.compile 引入了图捕获机制:

  • • TorchDynamo:在 CPython 字节码层面挂钩(利用 PyFrame_Eval / frame evaluation API),拦截 Python 执行帧,将符合追踪条件的 Tensor 操作记录为 FX Graph;遇到不可追踪的控制流则插入 graph break,退回 eager 执行。
  • • AOTAutograd:对捕获到的 forward 图做提前(ahead-of-time)自动微分,展开出 forward + backward 的联合 FX 图,此时的算子集合被规范化为 ATen 算子(aten:: 命名空间下的算子,约 2000+ 个)。

2.2 JAX/XLA 的路径:Trace-based Tracing

JAX 采用 jaxpr(JAX expression)作为中间表示,通过对 Python 函数插入抽象值(abstract tracer)实现符号执行式的图捕获,生成的 jaxpr 随后被下发给 XLA 编译为 HLO(High Level Optimizer IR)。

2.3 二者共性

无论是 FX Graph 还是 jaxpr,本质都是把动态图的 tensor 数据流"物化"为静态 DAG(有向无环图),节点是算子,边是张量依赖——这是后续所有转化的前提。


3. 算子分解:从"高层语义算子"到"原语算子"

框架里的一个算子(如 aten::layer_norm)往往不是硬件直接支持的原子操作,需要被分解(decomposition)为更基础的算子组合。

3.1 PyTorch 的分解表(Decomposition Table)

PyTorch 定义了两级算子集合:

  • • Core ATen IR:约 180 个核心算子,编译器后端只需实现这一层即可覆盖全部模型(其余算子通过分解表展开为核心算子的组合)。
  • • Prims IR(torch._prims):更细粒度的原语层,例如把 aten::softmax 展开为 exp、sum、div 的组合。

这种分层设计的意义在于:后端编译器不需要为每个高层语义算子单独写 lowering 规则,只需覆盖有限的原语集合,工程量大幅下降。

3.2 XLA 的 HLO 算子集

XLA 的 HLO 是一套更"数学化"的算子集合(dot, reduce, broadcast, convolution, custom-call 等),Transformer 中的 attention、matmul 等都会被表达为 dot_general + reduce 组合,再交由 XLA 的融合与布局分配 pass 处理。


4. 中间表示的分层降级(IR Lowering)

以 MLIR 生态(Torch-MLIR、IREE)为代表,算子转化呈现清晰的"方言(Dialect)"分层:

  • • Torch Dialect:贴近 PyTorch 语义,保留算子名字(如 torch.aten.matmul)。
  • • Linalg Dialect:将算子表达为结构化的循环嵌套 + 索引映射(linalg.generic 用 affine map 描述读写模式),这一层是矩阵乘法、卷积等算子能被自动向量化/并行化的关键——一个 linalg.matmul 本质上是对循环维度、tile 大小、内存布局的声明式描述,而非命令式代码。
  • • Affine/SCF:具体的循环结构(affine.for、scf.for),此时 tile 化、循环合并、向量化等 pass 开始作用。
  • • 最终降级到 LLVM IR(CPU 后端)或 NVVM/ROCDL(GPU 后端),再由 ptxas/llc 生成机器码。

这一路径与你在 Triton/ROCm 里接触到的 MFMA lowering 链条是同构的:Triton 的 TTIR → TTGIR → LLVM/AMDGPU IR 本质上也是这样一条"逐步具体化循环与内存访问模式"的降级链。


5. 算子融合:转化过程中的性能核心

如果不做融合,一个 LayerNorm + Softmax + Attention 序列会产生大量中间 tensor 的显存读写,成为访存瓶颈。融合是算子转化中最关键的优化环节。

5.1 融合的两个层次

  1. 1. 算子级融合(Operator Fusion):将 element-wise 算子链(如 add → gelu → dropout)合并为单个 kernel,减少 kernel 启动开销和显存往返。这类融合通常靠图重写规则(pattern rewrite)在 IR 层完成。
  2. 2. 循环级融合(Loop Fusion):在 Linalg/Affine 层,把多个 linalg.generic 的循环体合并,共享循环变量和内存访问,这是 TVM 的 Compute/Schedule 分离思想的核心——计算逻辑(compute)与执行调度(schedule,包括 tile、fuse、reorder、vectorize)是解耦的。

5.2 FlashAttention 作为"手工融合"的极致案例

标准 attention 会被分解为 matmul(Q,K) → softmax → matmul(P,V) 三个算子,产生 O(N²) 的中间矩阵。FlashAttention 本质上是一种算子融合的手工内核实现:通过分块(tiling)+ online softmax 重计算,将三个算子融合为一个 kernel,避免中间矩阵落地到 HBM。这也是为什么 FlashAttention 通常以 Triton/CUTLASS 手写 kernel 的形式存在,而不是靠编译器自动融合得到——目前通用编译器的融合 pass 还难以自动发现这种跨算子的数学等价重写(online softmax 的递推式)。


6. 算子到硬件指令的最终映射

6.1 Tensor Core / Matrix Core 层面

无论是 NVIDIA 的 mma/wgmma PTX 指令还是 AMD 的 MFMA 指令,编译器需要将 Linalg 层的 matmul 算子映射为特定的寄存器分片(fragment)布局:

  • • CUDA:nvcuda::wmma API 或直接的 mma.sync PTX 指令,涉及 warp 内线程对 A/B/C 矩阵片段的划分(对应你之前研究过的 MMA fragment mapping)。
  • • ROCm:Triton 的 third_party/amd/ 后端把 tt.dot 下降为 MFMA intrinsic,涉及 wavefront 内 64 个线程的数据重排(swizzle)以匹配 MFMA 指令的输入布局要求。

这一层的转化不再是通用编译器 pass,而是目标特定的 pattern matching + 寄存器分配,通常写死在后端代码里(如 Triton 的 ConvertTritonAMDGPUToLLVM)。

6.2 内存层次的映射

算子转化同时伴随着存储层级的分配决策:一个 linalg.generic 中的临时 buffer 会被决定放在寄存器、shared memory(LDS)还是全局显存,这由 Triton 的 tt.make_range/tt.load 的 block pointer 分析和 MLIR 的 bufferization pass 完成。这一决策直接决定了 kernel 是否会被 LDS bank conflict 或寄存器溢出(register spill)拖累。


7. 编译期 vs. 运行时算子转化

值得区分两种转化时机:

转化方式
代表框架
特点
AOT(提前编译)
XLA、TVM、torch.export + AOTInductor
编译期确定所有 shape 和算子选择,生成静态可执行文件,适合部署
JIT(即时编译)
TorchInductor(默认模式)、Triton autotune
运行时根据实际输入 shape 生成/缓存 kernel,支持动态 shape,但有首次编译开销

TorchInductor 是目前 PyTorch 默认后端,它将 ATen IR 转化为一种基于循环的 Python 中间表示(Loop-level IR),再自动生成 Triton kernel 代码(对 GPU)或 C++/OpenMP 代码(对 CPU)。这是"从算子到 Triton 源码"这一转化环节的具体实现——本质上是自动化了原本需要手工编写 Triton kernel 的过程。


8. 总结:转化链条的分层抽象

整个转化过程的核心工程哲学是分层抽象 + 逐步具体化:每一层只解决一类问题(图结构 → 算子语义 → 循环结构 → 内存布局 → 寄存器分配),使得编译器各层可以独立演进、独立测试。这与你在 QEMU/PCIe 虚拟化中"分层抽象设备行为"的工程思路是相通的——都是通过定义清晰的中间层接口,将一个复杂的端到端问题拆解为可组合、可替换的转化阶段。


TODO:深入研究 Linalg 的 tile/fuse pass 细节,Triton 的 MFMA lowering 

相关学习资料