运算流水线工作原理
2026/7/23大约 5 分钟
理解 PyTorch 代码如何一步步下沉到显卡晶体管进行暴算,核心在于看清一条“从人类抽象语义到机器物理指令”的流水线。
我们以大模型中最常用的一个简单操作 —— LayerNorm(层归一化) 为例,把 PyTorch、Triton、MLIR、LLVM 在这个过程中究竟在干嘛,用最直观的物理流水线给你连起来:
第一阶段:高级框架层(PyTorch)—— 提出高层意图
你在 Python 里写下这行代码:
out = torch.nn.functional.layer_norm(x, normalized_shape)
- PyTorch 在干嘛:它是一个高级调度中控。它本身不负责具体的矩阵运算,它只负责构建计算图、管理显存张量(Tensor)的生命周期。
- 传统的执行路径( eager 模式):PyTorch 看到这行代码,会直接通过 C++ 绑定(PyBind11),去调用英伟达官方写好的闭源加速库
cuDNN里的某个预编译好的layer_norm算子。 - 现代的编译路径(
torch.compile):PyTorch 2.0 之后,它不急着去调现成算子,而是用一个叫 TorchDynamo 的工具把这段 Python 代码抓取成一个静态计算图(FX Graph)。它发现后面紧跟着一个激活函数:out = torch.relu(out)。PyTorch 会大喊一声:“我想把这两个算子融合成一个,谁来帮我编排底层代码?” 于是,它把这个任务交给了 Triton。
第二阶段:算子表达层(Triton)—— 规划数据方块(Tile)的物理逻辑
Triton 在这里扮演的是高性能算子的高级生成器与编译器前端。
- Triton 在干嘛:传统的 CUDA C++ 需要工程师精细地去算“1号线程读1号显存”。而 Triton 引入了块级编程(Block-level)。
- 它的核心物理任务:Triton 接收到 PyTorch 的融合请求,自动或者由工程师手写一段 Triton Python 代码。这段代码的核心逻辑是:
- “从慢速主显存(HBM)里捞出一个形状为
[128]的数据方块(Tile)。” - “把这个方块塞进片上超高速缓存(SRAM)里,当场算完 Mean(均值)和 Variance(方差)。”
- “接着在 SRAM 里直接做完 ReLU 激活,最快速度把最终结果写回主显存。”
- 输出产物:Triton 将这段逻辑解析为高级语法树,但这时候它依然不知道具体怎么分配显卡的硬件寄存器,于是它把接力棒交给了 MLIR。
第三阶段:多级优化层(MLIR)—— 剥洋葱式的“精细化数据排布”
MLIR 是整个编译流水线里的超级智囊团。它最大的魔术在于方言(Dialect)系统,允许代码在不同的抽象层级上做剥洋葱式的优化。
- MLIR 在干嘛:它负责把 Triton 粗糙的“块级逻辑”一步步细化,转变成显卡硬件听得懂的并行操作。
- 它在后台演的这出接力赛:
- 第一步(Triton Dialect 层):MLIR 看着 Triton 丢过来的图,进行高层数学化简。它确认:“嗯,这是一个融合了 LayerNorm 和 ReLU 的连续数据块操作。”
- 第二步(Linalg / Vector Dialect 中间层):MLIR 开始进行硬核的并行排布和防坑优化。它计算出这个
[128]的方块该怎么分配给显卡的 32 个线程束(Warp)。为了防止这 32 个线程在读写片上 SRAM 时撞车,它会自动运行算法进行地址置换混淆(Swizzling)或申请空间填充(Padding),在底层默默帮你消灭 Bank 冲突。同时,它还会排布异步流水的时序,让计算单元在算当前块时,数据搬运单元悄悄去读下一块(实现 Overlap)。 - 第三步(LLVM / NVVM Dialect 底层):优化完访存和并行后,MLIR 把这些规则翻译成对应显卡厂牌(如英伟达)的专属底层硬件描述。
- 输出产物:经过 MLIR 层层抽丝剥茧、精细调优后的最底层 LLVM IR 代码。
第四阶段:通用芯片后端(LLVM)—— 临门一脚的“终极硬件打包”
LLVM 是整个工业界的硬件翻译大底座。
- LLVM 在干嘛:它不管什么是大模型,也不管什么是矩阵。它只负责把上一步生成的、极其规范的低级 LLVM IR 汇编,做最后的机器级物理映射。
- 它的核心任务:
- 精确计算并分配显卡内部极其珍贵的寄存器(Registers)空间。
- 将代码彻底翻译成特定显卡架构(比如英伟达 Hopper 架构)的物理指令。
- 输出产物:PTX(英伟达的高级汇编代码),并由显卡驱动进一步编译成 Cubin(物理二进制机器码)。
第五阶段:物理执行层(显卡硬件)—— 晶体管引爆算力
- 最终落地:最终生成的 Cubin 二进制机器码被注入 GPU 显存。
- 显卡的硬件调度器(Hardware Scheduler)瞬间拉起成千上万个物理线程,开始执行指令:数据从 HBM 疯狂抽向 SRAM,指令驱使 Tensor Core 的矩阵乘法器和 ALU 算术逻辑单元在高频时钟周期内进行电压翻转。
- 你在 PyTorch 里写的那行
layer_norm,在此刻真正变成了硅片上滚烫的电流和高频暴算。
💡 一句话总结它们究竟在干嘛:
PyTorch 负责出设计图纸(构建计算图),Triton 负责打好粗框架(规划分块与融合逻辑),MLIR 负责精细化施工优化(逐层优化访存、并行并消灭 Bank 冲突),LLVM 负责打包成商品交付(生成特定芯片的物理机器码)。
