From f4164eacd0f7b579aa49dbe98f4d4ee436cf7ff6 Mon Sep 17 00:00:00 2001 From: tlopex <820958424@qq.com> Date: Wed, 2 Sep 2026 01:54:01 -0400 Subject: [PATCH 1/2] Expand TIRx lowering pipeline documentation --- zh/tirx_guide/arch/index.rst | 1 + zh/tirx_guide/arch/lowering_pipeline.rst | 501 ++++++++++----- .../arch/tile_primitive_layout_lowering.rst | 590 ++++++++++++++++++ 3 files changed, 924 insertions(+), 168 deletions(-) create mode 100644 zh/tirx_guide/arch/tile_primitive_layout_lowering.rst diff --git a/zh/tirx_guide/arch/index.rst b/zh/tirx_guide/arch/index.rst index db9145de..c31d7a33 100644 --- a/zh/tirx_guide/arch/index.rst +++ b/zh/tirx_guide/arch/index.rst @@ -26,3 +26,4 @@ :maxdepth: 1 lowering_pipeline + tile_primitive_layout_lowering diff --git a/zh/tirx_guide/arch/lowering_pipeline.rst b/zh/tirx_guide/arch/lowering_pipeline.rst index d5ecd121..e542b031 100644 --- a/zh/tirx_guide/arch/lowering_pipeline.rst +++ b/zh/tirx_guide/arch/lowering_pipeline.rst @@ -15,161 +15,134 @@ specific language governing permissions and limitations under the License. +.. _chap_tirx_lowering_pipeline: + TIRx 编译流水线 =============== -``tvm.compile(mod, target, tir_pipeline="tirx")`` 接收一个 TIRx module,最终生成两部分代码:CPU 端的启动函数负责准备参数并启动 GPU,GPU 端的 kernel 负责执行计算。编译器不是一步完成这项转换,而是依次运行多个编译步骤。每个步骤称为一个 pass,负责对 IR 做一类特定的转换、检查或标注。 +调用 ``tvm.compile(mod, target, tir_pipeline="tirx")`` 时,输入并不是直接从 +Python 翻译成 PTX。编译器先消费 TIRx 特有的 tile primitive、execution scope +和 layout,再进行通用 TIR 的正规化、局部优化、类型合法化、host/device +拆分与代码生成。这里的 pass 是对 IR 执行一类转换、检查或标注的编译步骤; +target 指定设备和代码生成后端。 + +本页回答两个问题:**pipeline 里有哪些 passes,以及优化主要发生在哪里**。 +``Tx.gemm_async`` 如何使用 layout 推导 descriptor、直接 ``T.ptx.*`` 绕过什么, +单独放在 :ref:`chap_tirx_tile_layout_lowering` 中展开。 + +.. admonition:: 先给结论 + + 默认 TIRx 编译器可以概括为:**较薄的通用优化流水线,加上较厚的算子级 + lowering**。Kernel 作者或上层生成器负责 tile sizes、warp roles、同步、 + pipeline stages 和主要 layout;tile primitive 的实现负责合法性检查、 + descriptor/地址参数推导和指令分解;通用 passes 主要负责化简、合法化、 + 模块拆分和 ABI。 -TIRx 的完整 pass 顺序定义在 Apache TVM 源码中的 `python/tvm/tirx/compilation_pipeline.py -`_。 +.. note:: -整体编译路径 + 本书固定使用 Apache TVM ``0.26.0``。下面的 pass 名称和顺序都对应这个版本, + 并且讨论的是显式指定 ``tir_pipeline="tirx"`` 的路径。开发分支中的 pipeline + 可能已经变化;名为 ``"default"`` 的 pipeline 也不是 ``"tirx"`` 的别名。 + +完整编译路径 ------------ -``target`` 指定代码要在哪种硬件和后端上运行。下面的例子在 GPU 端使用 CUDA,在 CPU 端使用 LLVM:``tvm.compile`` 先将 target 信息写入 module,再运行模块级的 **tirx pipeline**。``tirx_pipeline`` 拆出 CPU 端的 host 函数和 GPU 端的 device 函数后,两者分别经过最后一轮面向具体 target 的转换,再交给相应的代码生成器: +先按职责观察整条路径: .. code-block:: text - 编写好的 TIRx - │ BindTarget - ▼ - tirx_pipeline - (SplitHostDevice 在其中拆分两条路径) - ├── host PrimFunc ──host finalization──▶ C/LLVM - └── device PrimFunc ─device finalization─▶ CUDA + TIRx Python + │ script parser + ▼ + TIRx PrimFunc + ├─ logical computation + ├─ Tx.* TilePrimitiveCall + ├─ Buffer + TileLayout + └─ device / CTA / warpgroup / warp / thread scopes + │ + │ BindTarget + ▼ + LowerTIRx + ├─ TilePrimitiveDispatch:Tx.* → target-specific TIR / T.ptx.* + └─ LowerTIRxCleanup:剩余 layout/access → physical buffer access + │ + ▼ + 结构正规化 + 局部程序变换 + dtype 合法化 + │ + ▼ + 校验 + SplitHostDevice + launcher ABI + ├─ host PrimFunc → host finalization → host backend + └─ device PrimFunc → device finalization → CUDA C++ / inline PTX + │ + ▼ + PTX / cubin / runtime module + +``PrimFunc`` 是 TIR 中的函数表示,finalization 是代码生成前面向具体 target +的最后一组转换。``BindTarget`` 在 ``tirx_pipeline`` 之前运行。Target 不只是决定最后使用哪个 +code generator;``TilePrimitiveDispatch`` 在较早阶段就需要 target,才能查找 +对应的算子实现。``SplitHostDevice`` 位于 module-level pipeline 后半段,拆分后 +host 和 device functions 才分别进入各自的 finalization。 + +高层 tile primitive 和直接 PTX 从不同位置进入这条路径: -``PrimFunc`` 是 TIR 中的函数表示。上图中的 host PrimFunc 就是 CPU 端的启动函数,device PrimFunc 则是 GPU 上执行的 kernel;``finalization`` 表示代码生成之前针对具体 target 所做的最后几步转换。 +.. code-block:: text -``tirx_pipeline`` 的 pass 顺序 -------------------------------- + Tx.gemm_async ── TilePrimitiveDispatch ──▶ T.ptx.tcgen05.mma ──┐ + ├─▶ 后续通用 passes + 直接 T.ptx.tcgen05.mma ────────────────────────────────────────┘ -下表按照实际执行顺序列出 ``tirx_pipeline`` 中的 19 个步骤。ABI 是函数之间的调用约定;表中的 ABI passes 负责把普通 TIR 函数改造成 runtime 能够调用的形式。``PassContext`` 是控制编译选项的配置对象:公共子表达式消除可以关闭,向量化和循环展开的行为也可以通过它调整。 +因此,直接 PTX 绕过的是对应 tile primitive 的算子级选择、检查和参数推导, +不是整个 TIRx pipeline。两条路径的详细比较见 +:ref:`chap_tirx_tile_layout_lowering`。 -.. list-table:: - :header-rows: 1 - :widths: 6 24 24 46 - - * - # - - 类别 - - Pass - - 作用 - * - 1 - - TIRx lowering - - ``LowerTIRx`` - - 完成 TIRx 的核心转换,详见下方 `LowerTIRx 的内部组成`_ - * - 2 - - TIR 规范化 - - ``UnifyThreadBinding`` - - 合并等价的 thread-axis bindings,使每个 ``threadIdx`` / ``blockIdx`` - axis 只声明一次 - * - 3 - - TIR 规范化 - - ``StmtSimplify`` - - 简化 IR 中的算术表达式 - * - 4 - - TIR 规范化 - - ``LowerTIRxOpaque`` - - 转换绑定到线程轴的循环,消除未标注的长度为 1 的循环,并规范化循环 pragma - * - 5 - - TIR 规范化 - - ``FlattenBuffer`` - - 将剩余的多维 TIR ``BufferLoad`` / ``BufferStore`` 展平为一维访问 - * - 6 - - 计算合法化 - - ``BF16ComputeLegalize`` - - target 不原生支持 ``bfloat16`` 计算时,将其提升到 ``float32`` 并改写为合法形式 - * - 7 - - TIR 规范化 - - ``NarrowDataType(32)`` - - 在能够证明安全时,将索引表达式和循环变量缩窄至 32 位 - * - 8 - - 循环转换 - - ``VectorizeLoop`` - - 将 ``T.vectorized`` 循环转换为向量操作;设置 ``tir.disable_vectorize`` 时,则改写为普通标量循环 - * - 9 - - 循环转换 - - ``UnrollLoop`` - - 展开标记为 ``T.unroll`` 的循环;普通常量循环只有在相应配置或 - pragma 启用时才会自动展开 - * - 10 - - TIR 规范化 - - ``StmtSimplify`` - - 向量化和循环展开后会出现更多可简化的常量,再次执行简化 - * - 11 - - TIR 规范化 - - ``CommonSubexprElim`` - - 将重复的子表达式提取为临时变量;设置 ``tir.disable_cse_tir`` 时跳过 - * - 12 - - 计算合法化 - - ``FP8ComputeLegalize`` - - target 不原生支持 ``float8`` 计算时,将其提升为受支持的类型 - (默认为 ``float32``) - * - 13 - - 校验与 ABI - - ``VerifyMemory`` - - 确保 host 代码不会直接解引用 device memory - * - 14 - - 校验与 ABI - - ``AnnotateEntryFunc`` - - 只有一个 PrimFunc 时直接将其标记为入口;有多个 PrimFunc 时,则标记其中唯一对外可见的函数 - * - 15 - - 校验与 ABI - - ``SplitHostDevice`` - - 识别 device regions,拆分 host 与 device PrimFuncs,并将 host 侧调用转换为 kernel-launch ABI - * - 16 - - 校验与 ABI - - ``LowerIket`` - - 普通 build 中移除 NVIDIA IKET annotations;启用 IKET 时则将其转换为 tracing 所需的形式 - * - 17 - - 校验与 ABI - - ``MakePackedAPI`` - - 将 host function 改写为 TVM runtime 通过 packed-function ABI 调用的形式 - * - 18 - - 存储合法化 - - ``FP8StorageLegalize`` - - target 不原生支持 ``float8`` storage 时,改用 ``uint8`` container 保存 - * - 19 - - 存储合法化 - - ``BF16StorageLegalize`` - - target 不原生支持 ``bfloat16`` storage 时,改用 ``uint16`` container 保存 +``LowerTIRx`` 的边界 +------------------------------ -Host 与 Device 的后续处理 -------------------------- +默认情况下,``LowerTIRx`` 包含两个 transformation passes: -上面列出的 19 个步骤组成 ``tirx_pipeline``。这条模块级 pipeline -结束后,``tvm.compile`` 会根据函数类型分别执行 finalization: +.. code-block:: text -- **host**:``LowerTVMBuiltin`` 处理 ``tvm_*`` builtins,``LowerIntrin`` - 处理面向具体 target 的 intrinsics。 -- **device**:``LowerWarpMemory`` 将 warp-scoped buffers 转换为 - shuffles,随后执行 ``StmtSimplify`` 和 ``LowerIntrin``。 + LowerTIRx = Sequential([ + TilePrimitiveDispatch, + LowerTIRxCleanup, + ]) -``LowerTIRx`` 的内部组成 ------------------------- +若设置 ``TVM_PRINT_AFTER_TIRX_DISPATCH_OPS``,两者之间还会临时插入一个 +``PrintIR``,但它只是调试 instrumentation,不是固定 transformation。 -``LowerTIRx`` 主要完成两个任务:为 tile-level 操作选择具体实现,以及把逻辑数据布局转换成实际的内存索引。它的核心转换由下面两个 passes 组成,定义在 Apache TVM 源码中的 -`src/tirx/transform/lower_tirx.cc -`_: +``TilePrimitiveDispatch`` 首先选择已注册的 target-specific 实现,把 +``Tx.copy``、``Tx.gemm_async``、``Tx.reduce`` 等 ``TilePrimitiveCall`` 替换为 +lower-level TIR 或 ``T.ptx.*``,并解析 device entry 内的 scope IDs。例如: .. code-block:: text - LowerTIRx = Sequential([ TilePrimitiveDispatch, LowerTIRxCleanup ]) + bx = T.cta_id([grid_x]) → bx = blockIdx.x + tx = T.thread_id([block_x]) → tx = threadIdx.x + +``LowerTIRxCleanup`` 随后对仍然存在的直接 ``BufferLoad`` / ``BufferStore`` 应用 +memory layout,展平相应 buffers,清除已消费的 layout metadata,并移除 +``tirx.buffer_offset`` wrappers。必须先 dispatch 后 cleanup:算子实现需要在 +metadata 消失前读取完整的 region、shape、dtype、scope 和 layout。 -- **``TilePrimitiveDispatch``** 为 tile 操作选择具体实现。TIRx 中的 ``copy``、 - ``gemm``、``reduction`` 等操作以 ``TilePrimitiveCall`` 表示;这个 pass 根据 - backend 选择对应实现。它还会把 ``T.cta_id``、``T.thread_id`` 等抽象的执行范围编号转换成 kernel launch 参数和线程绑定。 -- **``LowerTIRxCleanup``** 将逻辑坐标转换为物理索引。它把支持的逻辑 layout 应用到 buffer access 上,使后续 passes 可以直接处理具体的索引表达式。 +``LowerTIRx`` 成功结束后: -完成 ``LowerTIRx`` 后,tile 操作已经换成选定的底层实现,逻辑 layout 也已经落实为物理索引,``T.cta_id`` 和 ``T.thread_id`` 等抽象编号则变成了线程绑定。此时仍可能保留 thread-binding loops 和 TIRx 特有的 loop annotations;后续的 -``LowerTIRxOpaque`` 会规范化这些结构,再由 ``tirx.transform.FlattenBuffer`` -展平普通 TIR 中的 buffer access。 +- ``TilePrimitiveCall`` 已被具体实现替换; +- 抽象 scope IDs 已变成 launch parameters、``Bind`` 和 thread bindings; +- layout 已被算子 lowering 消费,或被 cleanup 物化为后端能处理的地址; +- ``T.ptx.*`` 等 target intrinsics 仍可存在; +- thread-binding loops 和 TIRx-specific loop annotations 仍可能存在,随后由 + ``LowerTIRxOpaque`` 规范化; +- host/device split、ABI lowering 和最终 code generation 尚未完成。 -一个简单 Kernel 的编译过程 --------------------------- +因此,这个阶段只能说 **TIRx 的核心高层语义已经消解**,不能说已经得到最终 +PTX,也不宜把它描述成编译流程的终点。算子 dispatch 和 layout cleanup 的内部 +过程见 :ref:`chap_tirx_tile_layout_lowering`。 -下面用一个 scale kernel 观察两件事:``T.cta_id`` 和 ``T.thread_id`` 怎样落实为具体的线程编号,以及一个 TIRx 函数如何拆成 CPU 端的启动函数与 GPU 上执行的 kernel。这个 kernel 处理 1,024 个元素,使用 4 个 CUDA thread blocks(CTA),每个 CTA 包含 256 个线程。 +最小示例:从 scope IDs 到 host/device +--------------------------------------- -**1. TIRx 源码使用抽象的线程编号。** +下面用一个不依赖 Tensor Core 的 scale kernel 串起这些阶段。它启动 4 个 CTAs, +每个 CTA 有 256 个 threads: .. code-block:: python @@ -183,37 +156,32 @@ Host 与 Device 的后续处理 T.device_entry() bx = T.cta_id([4]) tx = T.thread_id([256]) - B[bx * 256 + tx] = A[bx * 256 + tx] * T.float32(2.0) + i = bx * 256 + tx + B[i] = A[i] * T.float32(2.0) -``T.device_entry()`` 标记 GPU 代码的入口。``LowerTIRx`` 根据这个标记建立线程绑定;后面的 ``SplitHostDevice`` 再从所得设备代码区域(device region)中拆出单独的 device kernel。``T.cta_id([4])`` 表示 x 方向有 4 个 CTA, -``T.thread_id([256])`` 表示每个 CTA 有 256 个线程。这里的 ``bx`` 和 ``tx`` -仍是 TIRx 提供的抽象编号。 - -**2. ``LowerTIRx`` 将抽象编号转换为 TIR 线程绑定。** 它把 ``bx`` 和 ``tx`` -分别绑定到 ``blockIdx.x`` 与 ``threadIdx.x``。省略 buffer declarations 后,核心计算等价于: +``T.device_entry()`` 标出 device region。``LowerTIRx`` 将抽象 IDs 解析成 +``blockIdx.x`` 和 ``threadIdx.x``;省略 buffer declarations 后,核心结构可写成: .. code-block:: python with T.launch_thread("blockIdx.x", 4) as bx: tx = T.launch_thread("threadIdx.x", 256) - B[bx * 256 + tx] = A[bx * 256 + tx] * T.float32(2.0) - -这里仍然是 TIR,还不是 CUDA 源码。这段代码只保留了关键映射,不是编译器输出的完整 IR;下一节会给出打印完整结果的命令。 + i = bx * 256 + tx + B[i] = A[i] * T.float32(2.0) -**3. 后续 passes 拆分 host/device,并生成 CUDA。** 编译开始时只有一个 TIRx -函数。``LowerTIRx`` 生成线程绑定和 device region 后,``SplitHostDevice`` 将其拆成两个 TIR 函数(PrimFunc): +这仍是 TIR 的结构摘要,不是 CUDA 源码,也不是完整的 printer output。 +``SplitHostDevice`` 随后把原来的一个 ``PrimFunc`` 拆成两部分: .. code-block:: text - host launcher(由 scale 生成) - └── 启动 scale_kernel,gridDim.x = 4,blockDim.x = 256 + host launcher + └─ 启动 scale_kernel,gridDim.x = 4,blockDim.x = 256 device scale_kernel - └── 每个 GPU 线程将一个输入元素乘以 2 + └─ 每个 GPU thread 将一个元素乘以 2 -host 函数保存 kernel 的启动逻辑;device 函数保存真正的逐元素计算。 -``MakePackedAPI`` 随后将 host 函数转换为 TVM runtime 使用的统一调用形式。Device 函数则交给 CUDA backend,生成与下面代码等价的 CUDA -kernel: +``MakePackedAPI`` 把 host entry 降低为 runtime 使用的统一 ABI(函数调用约定)。 +Device function 则交给 CUDA backend,最终生成与下面代码等价的 kernel: .. code-block:: cuda @@ -222,41 +190,238 @@ kernel: B[i] = A[i] * 2.0f; } -整个过程可以概括为:TIRx 描述线程组织和计算,``LowerTIRx`` 将抽象编号落实为 -TIR 线程绑定,``SplitHostDevice`` 分开 CPU 端的启动逻辑与 GPU 端的计算,CUDA -backend 最后生成 CUDA 源码。 +后续 passes:按职责理解 +------------------------ + +``LowerTIRx`` 后的 passes 可以先分成四组;组内项目仍按源码定义的实际顺序 +执行,不能因为教学上分组就任意交换。``PassContext`` 是控制 pass 行为的 +编译配置对象。 + +.. list-table:: + :header-rows: 1 + :widths: 21 35 44 + + * - 职责 + - 主要 passes + - 做什么、不做什么 + * - 结构正规化 + - ``UnifyThreadBinding``、``StmtSimplify``、``LowerTIRxOpaque``、 + ``FlattenBuffer`` + - 统一 thread bindings,化简地址和条件,处理 thread-binding loops、 + unit loops 与 pragmas,并展平剩余 ``BufferLoad`` / ``BufferStore`` + * - 局部程序变换 + - ``NarrowDataType``、``VectorizeLoop``、``UnrollLoop``、第二次 + ``StmtSimplify``、``CommonSubexprElim`` + - 缩窄安全的 index;兑现已有 vectorize/unroll 标记或配置;执行局部代数 + 与公共子表达式化简,而不是自动发现完整 schedule + * - 类型合法化 + - BF16/FP8 compute legalization、BF16/FP8 storage legalization、 + final ``LowerIntrin`` + - 将当前 target/backend 不能直接表示的 compute、storage 和 intrinsic + 改写为受支持形式;部分 pass 可能因 target 能力而成为 no-op + * - 校验、模块和 ABI lowering + - ``VerifyMemory``、``AnnotateEntryFunc``、``SplitHostDevice``、 + ``LowerIket``、``MakePackedAPI`` + - 验证设备计算处于 thread environment,抽取 device kernel,降低 + kernel launch,并生成 runtime 可调用的 packed ABI + +``VectorizeLoop`` 主要降低已经写成 ``T.vectorized`` 的 loops;它不会自己搜索 +应该向量化哪一层;禁用 vectorization 时,相应循环按标量循环处理。 +``UnrollLoop`` 默认主要处理显式 ``T.unroll``,也可以通过 +``PassContext`` 或 pragma 设置自动展开阈值。两者都不等于自动 tiling 或 +software-pipeline synthesis。 + +Host/device split 与最终 codegen +-------------------------------- + +在 Apache TVM 0.26.0 中,``SplitHostDevice`` 是一个组合 pass:它识别 device +regions,将其抽取为 device ``PrimFunc``,并把 host 侧调用降低为 kernel-launch +约定。随后 ``MakePackedAPI`` 将公开的 host entry 改写为 TVM runtime 使用的 +packed-function ABI。 -这里不需要边界判断,因为 ``4 * 256`` 恰好等于 1,024。处理一般长度 ``N`` -时,需要向上取整得到 CTA 数量,并在 kernel 中判断 ``i < N``。 +.. code-block:: text -检查中间 IR 与生成代码 ----------------------- + 一个包含 device region 的 PrimFunc + │ + │ SplitHostDevice + ▼ + host launcher device PrimFunc + ├─ 准备调用参数 ├─ thread_extent + ├─ grid/block launch parameters ├─ physical buffer accesses + └─ 调用 device kernel └─ T.ptx.* / target intrinsics + │ │ + │ host finalization │ device finalization + ▼ ▼ + host target backend CUDA code generator + │ + ▼ + CUDA C++ source + inline PTX/helpers + │ + ▼ + CUDA frontend / assembler / runtime + +Module-level pipeline 结束后,finalization 分别运行: + +- **host**:``LowerTVMBuiltin``、``LowerIntrin``; +- **device**:``LowerWarpMemory``、``StmtSimplify``、``LowerIntrin``。 + +``LowerTVMBuiltin`` 处理 ``tvm_*`` builtins,``LowerIntrin`` 处理 +target-specific intrinsics。 +``LowerWarpMemory`` 将 warp-scoped buffers 降成 local storage 和 shuffle 等 +形式。CUDA code generator 随后生成 CUDA C++;部分 ``T.ptx.*`` intrinsic 会被 +打印为 inline PTX 或 helper code。之后 NVRTC/NVCC/ptxas 等 CUDA toolchain +组件才继续产生可加载的 PTX 或 binary。``LowerIntrin`` 本身并不等于“输出 PTX”。 + +编译器优化什么、不优化什么 +---------------------------- + +默认 TIRx pipeline 的责任边界如下: -为了查看中间 IR,可以只运行完整 pipeline 最前面的几步,然后停下来打印结果。下面先把 ``scale`` 以全局名 ``main`` 放入 ``IRModule``。CUDA target 指定 GPU -端生成 CUDA,``with_host("llvm")`` 则指定 CPU 端生成 LLVM 代码。 -``BindTarget`` 将这组 target 信息写入 module,随后只运行 ``LowerTIRx``: +.. list-table:: + :header-rows: 1 + :widths: 27 73 + + * - 层次 + - 主要责任 + * - Kernel 作者或生成 kernel 的 agent + - 选择 tile sizes、pipeline stages、warp roles、execution scope、同步、 + 主要 layout 和跨 tile 的整体 schedule + * - Layout helpers 与 tile dispatcher + - 根据显式 dtype/shape/mode 构造已知 layout;静态选择实现,检查合法性, + 推导 descriptor/物理参数,并将单个 tile operation 分解为硬件指令 + * - 通用 TIR passes + - 局部化简、index narrowing、显式或配置驱动的 vectorize/unroll,以及 + dtype、module 和 ABI 合法化 + * - CUDA backend 与 toolchain + - 生成 target code,进行更底层的 peephole、register allocation 和组装 + +所以,默认 pipeline 中确实有局部优化,但通常没有通用的自动 tiling、跨算子 +fusion、software-pipeline synthesis、warp-specialization synthesis、全局 +layout search、cost-model selection 或 autotuning。一个 dispatcher 根据 shape +将 tile operation 拆成多条硬件指令,是算子实现的一部分,不能与全程序 +schedule search 混为一谈。 + +检查完整流水线 +-------------- + +前面的 ``scale`` 例子也可以直接用于检查中间 IR。先绑定 CUDA device 与 LLVM +host target,再观察 ``LowerTIRx`` 前后的结果: .. code-block:: python + import tvm from tvm.tirx import transform as TT target = tvm.target.Target("cuda").with_host("llvm") mod = tvm.IRModule({"main": scale}) - mod = TT.BindTarget(target)(mod) - mod = TT.LowerTIRx()(mod) # 运行 LowerTIRx,转换抽象线程编号 - print(mod.script()) # 查看 LowerTIRx 之后的 IR + bound = TT.BindTarget(target)(mod) -输出中应该能看到 ``blockIdx.x`` 和 ``threadIdx.x`` 对应的线程绑定,而原来的 -``T.cta_id`` 与 ``T.thread_id`` 已经消失。 + print("=== authored TIRx ===") + print(bound.script()) -要查看最终 CUDA,可以运行完整 pipeline。这里的 host module 只导入了一个 -device module,因此 ``imports[0]`` 就是生成的 CUDA module; -``inspect_source()`` 返回它的源码: + print("=== after LowerTIRx ===") + lowered = TT.LowerTIRx()(bound) + print(lowered.script()) + +要检查完整 pipeline 生成的 CUDA source: .. code-block:: python - exe = tvm.compile(tvm.IRModule({"main": scale}), target=target, tir_pipeline="tirx") - cuda_mod = exe.mod.imports[0] - print(cuda_mod.inspect_source()) + exe = tvm.compile( + tvm.IRModule({"main": scale}), + target=target, + tir_pipeline="tirx", + ) + print(exe.mod.imports[0].inspect_source()) + +这个例子的 host module 只导入一个 device module,因此 ``imports[0]`` 就是 +CUDA module。如果问题发生在 ``Tx.*`` 到 ``T.ptx.*`` 之间,应进一步单独观察 +``TilePrimitiveDispatch``;Blackwell target 的设置和具体方法见 +:ref:`chap_tirx_tile_layout_lowering`。 + +Pass 顺序参考 +------------- + +Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时按下面 +顺序执行,共 19 步。若设置 ``tir.disable_cse_tir=True``,第 11 步会被省略, +后续 passes 依次前移。 + +.. list-table:: + :header-rows: 1 + :widths: 6 29 65 + + * - # + - Pass + - 作用 + * - 1 + - ``LowerTIRx`` + - Dispatch tile primitives、解析 execution scope,并物化 layout + * - 2 + - ``UnifyThreadBinding`` + - 合并等价的 thread-axis bindings + * - 3 + - ``StmtSimplify`` + - 使用 arithmetic analyzer 化简 statements 和索引表达式 + * - 4 + - ``LowerTIRxOpaque`` + - 转换 thread-binding loops、消除无 annotation 的 unit loops,并规范化 + loop pragmas + * - 5 + - ``FlattenBuffer`` + - 展平剩余 ``BufferLoad`` / ``BufferStore`` + * - 6 + - ``BF16ComputeLegalize`` + - target 不支持时将 BF16 compute 提升至 ``float32`` 并改写 + * - 7 + - ``NarrowDataType(32)`` + - 能够证明安全时,将 index/loop expressions 缩窄到 32 bits + * - 8 + - ``VectorizeLoop`` + - Lower 已标记的 vectorized loops;禁用时将其作为标量 loops 处理 + * - 9 + - ``UnrollLoop`` + - 展开显式 ``T.unroll``,并执行配置或 pragma 允许的自动展开 + * - 10 + - ``StmtSimplify`` + - 在 vectorize/unroll 暴露常量后再次化简 + * - 11 + - ``CommonSubexprElim`` + - 执行可配置关闭的公共子表达式消除 + * - 12 + - ``FP8ComputeLegalize`` + - target 不支持时将 FP8 compute 提升至默认的 ``float32`` 并改写 + * - 13 + - ``VerifyMemory`` + - 检查 GPU target/default calling-convention function 中,参数 buffer 的 + load/store 是否处于 ``thread_extent`` 环境 + * - 14 + - ``AnnotateEntryFunc`` + - 单一 PrimFunc 时直接标记;多函数 module 中标记唯一的公开 PrimFunc + * - 15 + - ``SplitHostDevice`` + - 标注/抽取 device functions,并 lowering host-to-device kernel calls + * - 16 + - ``LowerIket`` + - IKET 未启用时移除相关 annotations;启用时生成所需 tracing 形式 + * - 17 + - ``MakePackedAPI`` + - 将 host entry 改写为 runtime packed-function ABI + * - 18 + - ``FP8StorageLegalize`` + - target 需要 fallback 时,将 FP8 storage 改写为等宽 ``uint8`` 表示 + * - 19 + - ``BF16StorageLegalize`` + - target 需要 fallback 时,将 BF16 storage 改写为等宽 ``uint16`` 表示 + +不要仅凭 pass 名称推断默认会进行 aggressive auto-vectorization 或 +auto-unrolling;还需要检查 loop annotations 和当前 ``PassContext`` 配置。 + +版本与核心源码 +-------------- + +- `compilation_pipeline.py`_:module-level pipeline 与 host/device finalization; +- `lower_tirx.cc`_:``LowerTIRx`` 的两个 transformation passes 和调试 + ``PrintIR`` 插入点。 -生成的代码中应该能找到 ``blockIdx.x``、``threadIdx.x``,以及将每个输入元素乘以 2 的计算。 +.. _compilation_pipeline.py: https://github.com/apache/tvm/blob/v0.26.0/python/tvm/tirx/compilation_pipeline.py +.. _lower_tirx.cc: https://github.com/apache/tvm/blob/v0.26.0/src/tirx/transform/lower_tirx.cc diff --git a/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst b/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst new file mode 100644 index 00000000..09543c94 --- /dev/null +++ b/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst @@ -0,0 +1,590 @@ +.. Licensed to the Apache Software Foundation (ASF) under one + or more contributor license agreements. See the NOTICE file + distributed with this work for additional information + regarding copyright ownership. The ASF licenses this file + to you under the Apache License, Version 2.0 (the + "License"); you may not use this file except in compliance + with the License. You may obtain a copy of the License at + +.. http://www.apache.org/licenses/LICENSE-2.0 + +.. Unless required by applicable law or agreed to in writing, + software distributed under the License is distributed on an + "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + KIND, either express or implied. See the License for the + specific language governing permissions and limitations + under the License. + +.. _chap_tirx_tile_layout_lowering: + +Tile Primitive 与 Layout Lowering +================================= + +上一页 :ref:`chap_tirx_lowering_pipeline` 给出了完整 pipeline。本页只深入 +``LowerTIRx``,并沿一条 Blackwell ``Tx.gemm_async`` 解释 layout、算子静态分派 +和直接 PTX 之间的关系。 + +.. admonition:: 三个最容易混淆的结论 + + 1. 默认 row-major layout 的地址可能与传统连续数组完全相同,但 **“映射结果 + 相同”不等于“IR 中没有 layout metadata”**。 + 2. ``Tx.gemm_async`` 不靠 layout 执行矩阵乘法;它读取 layout 以匹配/验证 + operand 的物理约定,并推导 descriptor、offset 和硬件指令参数。 + 3. 直接 ``T.ptx.*`` 省掉的是对应 ``Tx.*`` 的算子级 lowering。周围的 scope + lowering、普通地址展开、合法化、host/device split 和 codegen 仍然存在。 + +贯穿示例:一条 ``Tx.gemm_async`` +-------------------------------- + +下面抽取 :ref:`chap_tirx_primer` 中单-tile GEMM 的关键部分。完整 kernel 还 +包含 SMEM/TMEM allocation、barrier 初始化、等待和释放;这里仅保留与 lowering +有关的声明与 tile operations: + +.. code-block:: python + + BLK_M, BLK_N, BLK_K = 128, 128, 64 + + A_layout = mma_shared_layout( + "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_M, BLK_K) + ) + B_layout = mma_shared_layout( + "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_N, BLK_K) + ) + + Asmem = pool.alloc((BLK_M, BLK_K), "float16", layout=A_layout) + Bsmem = pool.alloc((BLK_N, BLK_K), "float16", layout=B_layout) + + tmem = T.decl_buffer( + (128, 512), + "float32", + scope="tmem", + allocated_addr=tmem_addr[0], + layout=TileLayout(S[(128, 512) : (1@TLane, 1@TCol)]), + ) + + if warp_id == 0: + if T.ptx.elect_sync(): + Tx.gemm_async( + tmem[:, :BLK_N], Asmem[:, :], Bsmem[:, :], + accum=False, dispatch="tcgen05", cta_group=1, + ) + + Dreg = T.alloc_local((BLK_N,), "float32") + Dreg_wg = Dreg.view( + 128, + BLK_N, + layout=TileLayout(S[(128, BLK_N) : (1@tid_in_wg, 1)]), + ) + Tx.wg.copy_async(Dreg_wg[:, :], tmem[:, :BLK_N]) + +这些语句已经给 lowering 提供了大部分 schedule: + +.. list-table:: + :header-rows: 1 + :widths: 29 71 + + * - 输入信息 + - 它表达什么 + * - A/B/C 的 regions + - 本次 GEMM 的逻辑 ``M``、``N``、``K`` 范围 + * - A/B 的 SMEM layouts + - shared-memory 中的字节排列,以及选定的 128-byte swizzle + * - C 的 TMEM layout + - 声明期望的 accumulator datapath;dispatcher 可从特定 layout 推断 + ``.ws``,随后验证它与最终 tcgen05 datapath 一致并提取 slice offsets + * - ``warp_id`` 和 ``elect_sync`` + - 将 issuing scope 限定到一个被选中的 thread + * - ``dispatch="tcgen05"`` + - 强制选择 Blackwell tcgen05 variant + * - ``Dreg_wg`` layout + - 声明 readback 后每个逻辑元素的 thread ownership 和局部 slot + +其他 copy、reduce 等 tile primitive 也使用相同的 dispatcher 框架,但各自读取的 +layout 字段、约束和 lowering 结果并不相同。本例不能代表所有算子的具体硬件规则。 + +``TilePrimitiveDispatch`` 如何选择实现 +------------------------------------------------ + +前端把 ``Tx.gemm_async(...)`` 表示为 ``TilePrimitiveCall``。一次 dispatch 的输入 +由两部分合成: + +- ``TilePrimitiveCall`` 携带 operator、operand ``BufferRegion``、config 和 + ``dispatch=``; +- ``DispatchContext`` 提供 target、当前 execution scope、launch parameters、 + variable ranges 和插入初始化/分配语句所需的 callbacks。 + +候选实现按 ``(operator name, target kind)`` 注册,再按固定 priority 和 variant +名称排序。显式 ``dispatch="tcgen05"`` 只保留该 variant;没有显式指定时, +dispatcher 依次检查 predicates,并使用第一个成功返回 ``PrimFunc`` 的实现。 +这是 **静态规则选择**,不是 cost model,也不是 autotuning。 + +选中的 ``PrimFunc`` body 替换原 ``TilePrimitiveCall``。实现还可以通过 callbacks +请求 private allocation、device initialization、host initialization,或紧跟 +某个 buffer definition 的语句。因此,一次 lowering 不一定只在调用位置插入 +几条 PTX;它也可能准备 descriptor 或其他依赖资源。 + +这个 pass 还解析 ``T.device_entry()`` 内的抽象 scope IDs,生成 launch parameters、 +``Bind`` 和 ``thread_extent``。例如: + +.. code-block:: text + + T.cta_id([grid_x]) / T.thread_id([block_x]) + ↓ + blockIdx.x / threadIdx.x + +其中 ``grid_x`` 和 ``block_x`` 来自作者声明的 execution hierarchy,dispatcher +不会替 kernel 搜索 block size。 + +``Tx.gemm_async`` 的四步 lowering +--------------------------------------------- + +第一步:slice layout 并验证 operands +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Dispatcher 从三个 ``BufferRegion`` 取得 extents,并对各 buffer layout 执行 +``slice`` 与 ``canonicalize``。随后检查: + +- C 是否位于 TMEM,A/B 是否位于该 variant 支持的 memory scope; +- operand dtype 和逻辑 ``M/N/K`` 是否满足指令约束; +- A/B 的 sliced layout 是否包含受支持的 SMEM atom、swizzle 和 alignment; +- C 的 sliced layout 是否匹配受支持的 TMEM datapath。 + +这里不会把不兼容的 layout 自动“优化正确”。无法证明 layout、shape 或 +alignment 合法时,该 variant 会被拒绝;没有其他候选成功时,dispatch 失败。 + +第二步:A/B layout 变成 matrix descriptors +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +``tcgen05`` 从 SMEM 读取 A/B。Dispatcher 将 sliced layout 与硬件支持的 K-major +和 MN-major swizzle atoms 匹配,从匹配结果取得: + +.. code-block:: text + + swizzle mode + leading-dimension offset (ldo) + stride-dimension offset (sdo) + K-major / MN-major + 当前 MMA tile 相对 buffer 原点的 16-byte offset + +这些字段与 shared-memory base address 一起构成 matrix descriptor。前面的 +``Tx.cta.copy`` 或 ``Tx.copy_async`` 与后面的 MMA 因而通过同一个 layout contract +解释 SMEM 中的字节排列,kernel 作者不必手写 descriptor bit fields。 + +第三步:验证 C datapath 并取得 TMEM 目标 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +示例 C layout 的正向映射是: + +.. code-block:: text + + C[m, n] → { TLane: m, TCol: n } + +它声明期望怎样解释 accumulator 的硬件坐标。Dispatcher 确定最终 tcgen05 +instruction datapath,验证 C layout 与它相容,并提取 sliced region 的 +``TLane`` / ``TCol`` offsets;若 region 从非零 column 开始,slice 会产生相应的 +``TCol`` offset。最后再与 ``allocated_addr`` 组合成目标 TMEM address。 + +这里不能理解为“任意 C layout 都能改变硬件 accumulator 排列”。Layout-E +等受支持形式可以影响 ``.ws`` 推断,但最终仍必须匹配 tcgen05 能表达的 datapath。 + +第四步:shape、dtype 和 config 决定指令分解 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +``Tx.gemm_async`` 表示整个逻辑 tile,不等于“一行 DSL 固定对应一条 PTX”。 +Dispatcher 主要根据逻辑 ``M/N/K``、dtype、``cta_group`` 和 operator config 选择 +合法 instruction shape;已有 layout 用来验证 operand 约定并取得每一步的 +descriptor offsets。 + +对于示例中的 fp16 ``128×128×64`` GEMM,tcgen05 每个 K step 处理 16 个 fp16 +K elements,因此会产生 4 个 MMA iterations。去掉函数签名和大量常量细节后, +dispatch 后的结构可以概括为: + +.. code-block:: text + + # lowering sketch:尖括号内容不是可调用的 Python API + desc_a = <由 Asmem base、ldo、sdo、swizzle 组成的 TIR expression> + desc_b = <由 Bsmem base、ldo、sdo、swizzle 组成的 TIR expression> + desc_i = T.uint32() + + for ki in T.unroll(4): + T.ptx.tcgen05.mma( + , + , + , + desc_i, + enable_input_d=(ki != 0), + ... + ) + +``desc_i`` 对 dense tcgen05 路径是 dispatcher 在编译期间算出的 ``uint32`` +常量,不是 runtime 再调用某个 descriptor encoder。``T.unroll(4)`` 则由后面的 +``UnrollLoop`` 展开,随后的 ``StmtSimplify`` 会化简各次迭代的常量。由此可以看出: +instruction decomposition 来自 tcgen05 operator lowering,而 loop 展开属于 +后续通用 pass。 + +Readback 延续同一个物理约定 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +tcgen05 按最终确定的 datapath 把结果写入 TMEM;C layout 用于验证并解释这个 +物理结果。随后的 ``Tx.wg.copy_async`` 同时读取 C layout 和 ``Dreg_wg`` layout, +选择匹配的 ``tcgen05.ld`` form,并把逻辑 ``(m, n)`` 分配给 +``tid_in_wg=m`` 的 thread 及其局部 slot ``n``: + +.. code-block:: text + + GMEM + │ Tx.cta.copy / Tx.copy_async:按 A/B layout 写入 + ▼ + swizzled SMEM + │ Tx.gemm_async:descriptor 按同一 layout 读取 + ▼ + TMEM (TLane, TCol) + │ Tx.wg.copy_async:按 C 与 register layouts 解释和读取 + ▼ + per-thread local slots (tid_in_wg, m) + +Layout 没有执行 copy 或 MMA;它是 producer 与 consumer 对“同一个逻辑元素 +位于哪里”的共同约定。 + +Layout 的完整生命周期 +---------------------- + +同一个 layout 从创建到消失会经过下面几个阶段: + +.. code-block:: text + + parser 默认构造 / helper synthesis / 作者显式构造 + │ + ▼ + attach 到 Buffer + │ + ▼ + Buffer view / region slice + │ + ▼ + canonicalize / match + │ + ┌───────────┴───────────┐ + ▼ ▼ + operator dispatcher LowerTIRxCleanup + 消费硬件语义与 offsets 物化剩余的 memory offset + └───────────┬───────────┘ + ▼ + layout metadata 被清除 + │ + ▼ + backend 只看到地址和 intrinsics + +Dispatcher 和 cleanup 不是互斥的二选一。同一个 shared-memory swizzle 可以被 +GEMM dispatcher 用来构造 descriptor,同时 cleanup 仍会把这个 buffer 上残留的 +普通 ``BufferLoad`` / ``BufferStore`` 展开成物理地址。 + +.. list-table:: + :header-rows: 1 + :widths: 26 37 37 + + * - Layout 类型 + - Operator dispatcher 怎样使用 + - Cleanup 怎样使用 + * - 单一 memory axis ``m``,可含 swizzle + - 若 buffer 参与 tile primitive,可读取它以推导 descriptor、vector width + 或 offsets + - 将剩余直接 access 物化为一个线性 offset + * - ``laneid`` / ``tid_in_wg`` 加局部 ``m`` + - Register-aware operator 将它解释为 thread ownership 与局部 slot + - 不能直接遗留;必须先通过理解它的 operator,或用 ``.view()`` 后再以 + ``.local()`` 取得当前 thread 的 storage view + * - ``TLane`` / ``TCol`` + - TMEM-aware operator 验证 datapath 并取得硬件地址 + - 不能被普通 TIR ``BufferLoad`` 压成一个 pointer offset + +默认 row-major 不等于“没有 layout” +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +在本章讨论的 CUDA TIRx ``PrimFunc`` 中,``T.match_buffer``、``T.decl_buffer`` +等 buffer APIs 的 ``layout`` 参数默认值就是 ``"default"``。因此作者省略 +``layout=`` 时,parser 会自动构造: + +.. code-block:: text + + TileLayout(S[shape]) + +它是 dense row-major layout。三种写法的差别如下: + +.. list-table:: + :header-rows: 1 + :widths: 24 35 41 + + * - 写法 + - 普通连续 access 的结果 + - IR 中保留的信息 + * - 省略 ``layout=``,或写 ``layout="default"`` + - cleanup 后与传统 row-major 地址相同 + - 有一个可供 slice、match 和 operator 检查的 ``TileLayout`` + * - ``layout=None`` + - 通过普通 shape/stride 规则也可得到相同地址 + - 明确不附带 layout metadata;operator 无法从它取得专用映射 + * - 显式硬件 layout + - 地址可能包含 padding/swizzle,或映射到 thread/TMEM axes + - 携带 operator-specific 的物理约定 + +所以,对普通连续 buffer 来说,“默认 layout”和“没有 layout”可能产生完全 +相同的最终地址;差别在于编译器前半程有没有一份统一、可检查的映射契约。 +它本身不会凭空带来性能收益,更不等于编译器已经选出了最佳 layout。 + +Layout 自动到什么程度 +~~~~~~~~~~~~~~~~~~~~~~ + +“自动 layout”常混用三种含义: + +1. **默认构造。** Parser 在省略参数时补 dense row-major layout。这只是默认 + 语义,不是性能搜索。 +2. **Helper synthesis。** ``mma_shared_layout``、``tmem_datapath_layout``、 + ``tcgen05_atom_layout`` 等 helper 根据显式 dtype、shape 和 mode 构造已知 + 硬件 layout。选择哪个 helper 和 mode 仍由作者或上层生成器决定。 +3. **Lowering-time parameter inference。** Dispatcher 结合既有 layout 与 + shape、dtype、``cta_group`` 和 config,推导 major mode、descriptor fields、 + instruction shape 与 offsets。它是从已给定约定推出底层参数,不是反向搜索 + 最优 layout。 + +默认 TIRx pipeline 没有一个全局 ``InferOptimalLayout`` pass,也不会从任意 +手写 PTX 反推出 lane/register、SMEM swizzle 或 TMEM layout。 + +三种 layout 怎样变成物理位置 +----------------------------- + +普通 memory layout:逻辑下标变成地址 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +考虑逻辑 shape ``(4, 8)``、row stride 为 16 的 padded layout: + +.. code-block:: text + + layout = TileLayout(S[(4, 8) : (16, 1)]) + + layout.apply(i, j)["m"] = i * 16 + j + + B[i, j] → B_flat[i * 16 + j] + +这是纯映射示例;实际 backing allocation 必须至少容纳 +``layout.span() = 56`` 个 elements,而不是只分配逻辑元素数 ``4 * 8 = 32``。 +若通过函数参数传入 B,caller 也必须满足这个物理容量约定。 + +``LowerTIRxCleanup`` 会先把 layout 产生的物理坐标写入 access index,并保留 +buffer 的 ``elem_offset`` metadata;后续 ``FlattenBuffer`` 再把 +``elem_offset`` 折入最终线性 index。所以上图表示的是最终 **有效地址语义**, +不是声称 cleanup 内某一个 AST 节点已经完成所有后续 folding。 + +如果使用 ``ComposeLayout``,物理 offset 还可能包含 shared-memory swizzle 的 +XOR、shift 和 mask。完整 layout 代数见 :ref:`chap_tirx_layout_api`。 + +Distributed layout:逻辑元素变成 ownership 与局部 slot +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Register-backed tile 通常需要两个物理坐标:哪个 thread 持有元素,以及它在该 +thread 的第几个局部 slot。例如: + +.. code-block:: python + + fragment_layout = TileLayout( + S[(8, 4, 2) : (4@laneid, 1@laneid, 1)] + ) + +把它解释成逻辑 ``8×8`` tile 时: + +.. code-block:: text + + laneid = 4 * row + col // 2 + m = col % 2 + +``laneid`` 表示 ownership,``m`` 表示 lane-local slot;``m`` 不是最终 PTX 中 +某个固定寄存器编号。真实 register allocation 仍由 CUDA toolchain 完成。 + +普通 TIR ``BufferLoad`` 无法只凭一个 offset 验证“当前 thread 是否拥有这个 +逻辑元素”。因此,含 thread axis 的 layout 如果直接遗留到 cleanup 会报错。 +``.view()`` 先建立 distributed logical view,随后还必须通过 ``.local()`` 取得 +当前 thread 的 storage view;``.view()`` 单独使用并不能让直接 load 合法。 + +TMEM layout:逻辑元素变成 ``TLane`` 与 ``TCol`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Blackwell TMEM 使用二维硬件坐标: + +.. code-block:: text + + TileLayout(S[(128, N) : (1@TLane, 1@TCol)]) + + C[m, n] → { TLane: m, TCol: n } + +``TLane`` 是 TMEM 的物理 lane row,不是执行当前代码的 CUDA ``lane_id``。 +``TLane`` 与 ``TCol`` 也不是普通 pointer 的两个 strides;tcgen05-aware operator +必须先验证并解释它们。若这种二维坐标直接留给普通 ``BufferLoad``,cleanup +无法把它降成所要求的单一 memory offset。 + +``LowerTIRxCleanup`` 的准确边界 +--------------------------------------------- + +Dispatcher 完成后,``LowerTIRxCleanup`` 运行 ``LayoutApplier``。对剩余普通 +``BufferLoad`` / ``BufferStore``,其核心工作是: + +.. code-block:: text + + logical indices + │ + ▼ + layout.canonicalize().apply(indices, shape) + │ + ▼ + 一个 symbolic physical coordinate + │ + ▼ + flattened buffer access + +对于 CUDA 的直接 memory access,有 layout 时最终必须只产生一个 physical +coordinate,通常命名为 ``m``;没有 layout 时则使用普通 shape/stride 规则。 +``LayoutApplier`` 还把 layout-backed buffers 重建为 physical views并清空 layout +metadata。随后 ``BufferOffsetRemover`` 消除 ``tirx.buffer_offset(BufferLoad)`` +wrapper,使其中的 offset 与已经展平的访问一致。 + +这里的 ``symbolic`` 表示编译器在 lowering 时构造、化简地址表达式;表达式中仍可 +包含 runtime loop index、``threadIdx`` 或函数参数,并非所有地址都在编译期变成 +常量。 + +完整 ``LowerTIRx`` 成功后,可以依赖: + +- 所有 ``TilePrimitiveCall`` 已被具体实现替换,否则 pass 会失败; +- scope IDs 已被解析; +- layout metadata 已被 operator 消费或由 cleanup 物化并清除; +- ``T.ptx.*`` 可以继续存在,尚未变成最终 CUDA source/PTX assembly; +- 类型合法化、host/device split 和 ABI lowering 仍未完成。 + +高层 tile primitive 与直接 PTX +------------------------------- + +``T.ptx.tcgen05.mma`` 是 target intrinsic ``Call``,不是 ``TilePrimitiveCall``。 +它不会进入 ``gemm_async`` 的 registered variants,但 surrounding kernel +仍会经过 scope lowering、cleanup、通用 passes、host/device split 和 codegen。 + +.. list-table:: + :header-rows: 1 + :widths: 30 35 35 + + * - 责任 + - ``Tx.gemm_async`` + - 直接 ``T.ptx.tcgen05.mma`` + * - Backend variant + - Dispatcher 选择,或验证显式 ``dispatch=`` + - 作者已经选择 + * - Logical shape 到 instruction tiles + - Dispatcher 根据 shape/dtype/config 分解 + - 作者手写每条 instruction + * - Operand layout 合法性 + - Dispatcher 做 operator-specific matching 和检查 + - 主要由作者保证 + * - SMEM/TMEM descriptor 与 offsets + - 从 layout、region 和 config 推导 + - 作者手写或调用低级 encoding intrinsics + * - Lane/local-slot/TMEM operand 顺序 + - Layout 与 dispatcher 共同表达和检查 + - 作者编码在 lane 公式、地址与 operand 顺序中 + * - Simplify、legalize、split、codegen + - 仍然执行 + - 仍然执行 + +“作者编码寄存器顺序”不是说作者指定 ``%r17`` 这样的最终物理寄存器号;作者 +只是在 local arrays、lane formulas 和 PTX operands 中表达值的相对位置,真正的 +register allocation 仍由 ptxas 等工具完成。区别在于:直接 PTX 的 IR 不再保留 +“这个 operand 对应逻辑 ``C[m,n]``”的完整 tile-op contract,dispatcher 因而无法 +替作者做同等级的结构匹配与 layout mismatch 诊断。 + +直接 PTX 也不必然绕过所有 memory layout: + +- 若 intrinsic 参数经 ``buf.ptr_to(logical_indices)`` 构造,且 buffer layout + 能降低为 **单一 memory-axis offset**,cleanup 仍会把逻辑 indices 映射成 + 物理地址。 +- Distributed/thread-axis layout 或 ``TLane/TCol`` layout 不能把 ``ptr_to`` + 当作通用逃生口;若它们仍以直接 ``BufferLoad`` 形式出现,cleanup 会报错。 +- 一旦 intrinsic 只接收 raw base address 和作者已经算好的 offset,layout + mapping 就已被绕过。``buf.data + raw_offset`` 是常见写法,但不是唯一方式。 + +因此,“直接 PTX 没有消灭 layout”的准确含义是:它只会让 **那个低级 +instruction call** 不再经过 tile-op layout inference;周围 buffers 的普通地址 +访问仍可能需要 layout。反过来,如果所有相关地址、lane mapping 和 operand +顺序都由作者以 raw expressions 写完,那么 layout 对这条指令当然不会再提供 +额外推导。 + +失败表示检查,而不是自动修复 +------------------------------ + +假设 C 位于 TMEM,却只给它一个普通 row-major ``m`` layout,然后强制 +``dispatch="tcgen05"``。该 layout 无法证明逻辑 C region 与 ``TLane/TCol`` +datapath 相容,tcgen05 implementation 会拒绝它;如果没有其他候选,编译器 +报告 dispatch failure。 + +这类错误说明了 TIRx 的责任边界:dispatcher 会从 **已经声明的** layout 推导 +底层参数并验证约束,但不会搜索一个新 layout,再悄悄重写 kernel 的 producer、 +consumer 和 allocation 使它们全部匹配。 + +检查 dispatch 与 layout lowering +-------------------------------- + +调试时可以在 cleanup 删除 layout metadata 前,单独查看 dispatch 结果: + +.. code-block:: python + + import tvm + from tvm.tirx import transform as TT + + target = tvm.target.Target("cuda -arch=sm_100a").with_host("llvm") + mod = tvm.IRModule({"main": kernel}) + bound = TT.BindTarget(target)(mod) + + print("=== authored TIRx ===") + print(bound.script()) + + print("=== after TilePrimitiveDispatch ===") + dispatched = TT.TilePrimitiveDispatch()(bound) + print(dispatched.script()) + + print("=== after LowerTIRx ===") + lowered = TT.LowerTIRx()(bound) + print(lowered.script()) + +``TilePrimitiveDispatch`` 运行时必须能解析出 target;显式 ``BindTarget`` 最清晰, +也可以依赖已有的 PrimFunc target attribute 或 current target context。这里指定 +``sm_100a`` 是因为示例使用 Blackwell ``tcgen05``;生成和运行最终代码还要求相应 +CUDA toolkit 与硬件支持。 + +``LowerTIRx`` 也提供一个调试开关,在 dispatch 与 cleanup 之间打印 IR: + +.. code-block:: bash + + TVM_PRINT_AFTER_TIRX_DISPATCH_OPS=1 python your_kernel.py + +检查 ``Tx.gemm_async`` 时,建议依次确认: + +1. 原始 C/A/B regions 和 layouts; +2. dispatch 选择了哪个 variant; +3. A/B descriptors 的 major、swizzle、``ldo/sdo`` 和 slice offsets; +4. C datapath 与 TMEM offsets; +5. 生成的 MMA iteration 数量与 ``enable_input_d``; +6. cleanup 后还剩哪些直接 physical accesses。 + +这样可以把问题定位到 kernel schedule、layout contract、operator dispatcher +或后端 codegen,而不是只比较 TIRx 源码和最终 assembly。 + +核心源码导航 +------------ + +- `dispatcher.py`_:variant registry、priority、predicate 与失败报告; +- `tile_primitive_dispatch.cc`_:scope/launch context、body replacement 与 + callbacks; +- `lower_tirx_cleanup.cc`_:``LayoutApplier``、``BufferOffsetRemover`` 与 + physical-offset materialization; +- `tcgen05 gemm dispatcher`_:从 operand regions/layouts 推导 descriptors、 + TMEM address、instruction tiling 和 ``T.ptx.tcgen05.mma``。 + +.. _dispatcher.py: https://github.com/apache/tvm/blob/v0.26.0/python/tvm/tirx/operator/tile_primitive/dispatcher.py +.. _tile_primitive_dispatch.cc: https://github.com/apache/tvm/blob/v0.26.0/src/tirx/transform/tile_primitive_dispatch.cc +.. _lower_tirx_cleanup.cc: https://github.com/apache/tvm/blob/v0.26.0/src/tirx/transform/lower_tirx_cleanup.cc +.. _tcgen05 gemm dispatcher: https://github.com/apache/tvm/blob/v0.26.0/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py From 3301bfdacdaf978f4427110f76d8186cabe5f61f Mon Sep 17 00:00:00 2001 From: tlopex <820958424@qq.com> Date: Thu, 3 Sep 2026 18:56:00 -0400 Subject: [PATCH 2/2] Refine TIRx compiler internals documentation --- zh/tirx_guide/arch/index.rst | 3 +- zh/tirx_guide/arch/ir_representation.rst | 272 +++++++ zh/tirx_guide/arch/lowering_pipeline.rst | 202 ++--- .../arch/tile_primitive_layout_lowering.rst | 722 ++++++++---------- 4 files changed, 646 insertions(+), 553 deletions(-) create mode 100644 zh/tirx_guide/arch/ir_representation.rst diff --git a/zh/tirx_guide/arch/index.rst b/zh/tirx_guide/arch/index.rst index c31d7a33..f2ff1fc6 100644 --- a/zh/tirx_guide/arch/index.rst +++ b/zh/tirx_guide/arch/index.rst @@ -20,10 +20,11 @@ 编译器内部机制 ============== -本节介绍 TIRx 编译器如何将编写好的 module 转换成 CPU 端的启动函数和 GPU 端的 device code,并沿着编译流水线说明 TIRx 高层结构、host/device 拆分以及 CUDA 代码生成之间的关系。 +本节先说明 TIRx IR 怎样组织函数体、buffer、layout 和执行层级,再沿编译流水线观察这些信息怎样变成 CPU 端的启动函数和 GPU 端的 device code。 .. toctree:: :maxdepth: 1 + ir_representation lowering_pipeline tile_primitive_layout_lowering diff --git a/zh/tirx_guide/arch/ir_representation.rst b/zh/tirx_guide/arch/ir_representation.rst new file mode 100644 index 00000000..d83e046e --- /dev/null +++ b/zh/tirx_guide/arch/ir_representation.rst @@ -0,0 +1,272 @@ +.. Licensed to the Apache Software Foundation (ASF) under one + or more contributor license agreements. See the NOTICE file + distributed with this work for additional information + regarding copyright ownership. The ASF licenses this file + to you under the Apache License, Version 2.0 (the + "License"); you may not use this file except in compliance + with the License. You may obtain a copy of the License at + +.. http://www.apache.org/licenses/LICENSE-2.0 + +.. Unless required by applicable law or agreed to in writing, + software distributed under the License is distributed on an + "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + KIND, either express or implied. See the License for the + specific language governing permissions and limitations + under the License. + +.. _chap_tirx_ir_representation: + +TIRx IR 的组织方式 +================== + +:ref:`chap_tirx_primer` 已经展示了 scope、layout 和 dispatch 在 kernel 中的写法。本页选取其中一条 tile primitive,展开 parser 生成的对象,观察每类信息落在哪个节点、节点之间如何引用,以及这些节点为什么需要分开表示。 + +``@Tx.prim_func`` 解析后返回 ``tvm.tirx.PrimFunc``。它以 ``Stmt`` 树保存程序结构,以 ``PrimExpr`` 保存索引和标量计算;``Buffer``、``Layout`` 与 ``ExecScope`` 等对象由树中的节点引用。进入模块级 pass 前,一个或多个 ``PrimFunc`` 会作为全局函数放入 ``IRModule``。 + +通用基础设施与 TIRx 方言 +------------------------- + +TIRx 使用 TVM 的 ``IRModule``、``BaseFunc``、表达式基类、类型系统、``Op`` 注册表和 pass 管理机制。函数和语句结构采用 TIRx 方言节点,``IntImm``、``FloatImm`` 与 ``Range`` 等叶子对象来自通用 IR: + +.. code-block:: text + + TVM 通用基础设施 + ├─ IRModule + ├─ BaseFunc + ├─ ir.PrimExpr / IntImm / FloatImm / Range + ├─ Op / Type / Target + └─ PassContext 与 pass manager + │ + ▼ + TIRx 方言 + ├─ tirx.PrimFunc + ├─ tirx.Stmt + ├─ tirx.Var / tirx.Call / tirx.BufferLoad + ├─ tirx.Buffer / tirx.BufferRegion + ├─ tirx.ScopeIdDefStmt / tirx.ExecScope + └─ tirx.TilePrimitiveCall + +模块基础设施通过 ``PrimFunc`` 的运行时类型区分 TIRx 函数,函数体则由 TIRx 的 visitor 或 mutator 遍历。TIRx 由此接入 TVM 的通用模块,同时为 tile、layout 和执行层级保留专门的节点类型。 + +沿一条 tile primitive 看对象关系 +--------------------------------- + +以入门章计算阶段的 ``gemm_async`` tile primitive 和紧随其后的 ``tcgen05.commit`` intrinsic 为观察点。第一条调用在 parser 之后成为 ``TilePrimitiveCall``;第二条已经确定为 PTX 操作,保存为 ``Evaluate(Call)``。省略外围条件和其他语句后,局部对象关系如下: + +.. code-block:: text + + tirx.PrimFunc.body + └─ ... + ├─ TilePrimitiveCall + │ ├─ op = Op("tirx.tile.gemm_async") + │ ├─ args[0:3] + │ │ ├─ BufferRegion(tmem) + │ │ ├─ BufferRegion(Asmem) + │ │ └─ BufferRegion(Bsmem) + │ ├─ args[3:6] = transpose_A, transpose_B, accum + │ ├─ config = {"cta_group": 1} + │ ├─ dispatch = "tcgen05" + │ └─ scope = ExecScope("thread") + └─ Evaluate + └─ Call(op=tirx.ptx.*) + +这张局部对象图包含两种关系。``Stmt`` 节点之间的包含关系确定执行顺序和控制流;同一个 ``Buffer`` 则可以被声明节点、``BufferRegion`` 和访问节点共同引用。图中的箭头表示共享同一个 ``ObjectRef``,pass 根据对象身份连接定义与使用,再沿这些引用读取 layout 和 scope 等信息。 + +``PrimFunc`` 给出函数边界 +------------------------- + +``tirx.PrimFunc`` 的核心字段是 ``params``、``buffer_map``、``body`` 和 ``ret_type``,函数级 ``attrs`` 来自 ``BaseFunc``。其中 ``params`` 保存 handle 与标量参数,``body`` 保存一条 ``Stmt``;多条顶层语句由 ``SeqStmt`` 组织。 + +Buffer 参数在函数边界上分成两部分。参数注解 ``A: Tx.Buffer((M, K), "float16")`` 会在 ``params`` 中产生一个 handle 变量,同时由 ``buffer_map`` 将这个 handle 映射到包含 shape、dtype 和存储域的 ``Buffer`` 对象。分析 pass 可以直接从 ``buffer_map`` 读取参数约束,无需从函数体中的声明重新恢复。 + +``Tx.device_entry()`` 在 ``body`` 中形成 ``AttrStmt(attr_key="tirx.device_entry", value=True, body=...)``。它以一段语句区域为边界,因此保存在函数体内。 + +声明节点为何平铺在 ``SeqStmt`` 中 +---------------------------------- + +在各自所在的词法作用域内,变量绑定、buffer 声明和执行 ID 声明采用平铺形式,直接成为 ``SeqStmt`` 的成员: + +.. code-block:: text + + SeqStmt + ├─ ScopeIdDefStmt(tx = cta→thread, extent=128) + ├─ AllocBuffer(S) + ├─ Bind(i, ...) + ├─ TilePrimitiveCall(... S ...) + └─ BufferStore(... i ...) + +``Bind``、``AllocBuffer``、``DeclBuffer`` 和 ``ScopeIdDefStmt`` 都是独立的 ``Stmt``。这些节点分别保存绑定、分配或声明信息,外围 ``SeqStmt`` 负责承载后续程序。声明产生的变量或 buffer 对该 ``SeqStmt`` 中的后续语句可见,pass 按顺序遍历时便能维护当前可用的变量、buffer 和执行上下文。 + +控制流本身仍按词法作用域嵌套,索引、条件和标量计算以 ``PrimExpr`` 嵌入各个语句。这样的主干兼顾了控制流结构与声明的源码顺序,普通循环、表达式和 buffer 访问也能继续使用统一的 visitor 与 mutator。 + +分配与视图共享底层存储 +------------------------ + +``AllocBuffer`` 用一个 ``Buffer`` 描述新分配的存储,``DeclBuffer`` 则可以在已有数据指针上声明新的视图。入门 GEMM 中的 ``Dreg = Tx.alloc_local(...)`` 和 ``Dreg_wg = Dreg.view(...)`` 在 IR 中形成下面的关系: + +.. code-block:: text + + AllocBuffer + └─ buffer = Dreg + ├─ data ──────────────┐ + ├─ shape = [BLK_N] │ 同一个 data Var + └─ layout = ... │ + │ + DeclBuffer │ + └─ buffer = Dreg_wg │ + ├─ data ──────────────┘ + ├─ shape = [128, BLK_N] + └─ layout = distributed layout + +``AllocBuffer(Dreg)`` 记录每个 thread 的局部存储。``Dreg_wg`` 是另一个 ``Buffer`` 对象,由 ``DeclBuffer`` 加入语句序列;它保存自己的 shape 和 layout,同时与 ``Dreg`` 指向同一个 data ``Var``,dtype、strides 与 ``elem_offset`` 等属性从 ``Dreg`` 延续。入门示例中的 ``(128, BLK_N)`` 因而描述 warpgroup 使用的逻辑坐标系;在 parser 生成的 IR 中,实际分配节点仍是形状为 ``(BLK_N,)`` 的 ``AllocBuffer(Dreg)``。 + +``BufferRegion`` 连接 tile 与数据布局 +------------------------------------- + +``Buffer`` 是一份结构化数据视图,保存数据指针、dtype、shape、strides 和 ``elem_offset``,数据指针的类型携带 storage scope。TIRx 还在其中保存可选的 ``layout``,以及 TMEM 等专用存储使用的 ``allocated_addr``。参数的 ``buffer_map``、局部分配、数据访问和 tile 区域可以共同引用同一个 ``Buffer``。 + +单点访问形成 ``BufferLoad`` 或 ``BufferStore``,切片形成 ``BufferRegion``。``BufferRegion`` 只保存原 ``Buffer`` 和每一维的 ``Range(min, extent)``: + +.. code-block:: text + + TilePrimitiveCall.args[i] + │ + ▼ + BufferRegion + ├─ region = [Range(...), Range(...)] + └─ buffer ────────────────┐ + ▼ + Buffer + ├─ dtype / shape / storage scope + ├─ layout + └─ allocated_addr + +一条 tile primitive 因而可以同时引用本次操作覆盖的逻辑范围,以及底层数据视图中的 dtype、存储域、layout 和硬件地址。Layout 保存在 ``Buffer`` 上,因此各个 ``BufferRegion`` 都能访问同一份映射。 + +``layout`` 字段保存一个可选的 ``Layout`` 对象。``Layout`` 是统一的映射接口,当前有三种具体节点: + +.. code-block:: text + + Buffer.layout: Optional + │ + ▼ + Layout + ├─ TileLayout + │ ├─ shard: [Iter(extent, stride, Axis), ...] + │ ├─ replica: [Iter(extent, stride, Axis), ...] + │ └─ offset: {Axis: PrimExpr} + ├─ SwizzleLayout + │ ├─ per_element + │ ├─ swizzle_len / atom_len + │ └─ swizzle_inner + └─ ComposeLayout + ├─ swizzle: SwizzleLayout + └─ tile_layout: TileLayout + +``TileLayout`` 用一组 ``Iter`` 表示普通存储或线程映射,每个 ``Iter`` 保存 extent、stride 和物理 ``Axis``;``SwizzleLayout`` 保存 XOR swizzle 的参数;``ComposeLayout`` 同时持有一项 swizzle 和一项 tile mapping。三者通过 ``Layout`` 接口提供 ``Apply``、``Slice`` 和 ``Canonicalize`` 等操作。 + +入门 GEMM 的 ``128×64`` A/B shared-memory layout 在 canonicalize 后是 ``SwizzleLayout``,TMEM accumulator 使用 ``TileLayout``;带有额外外层 tile mapping 的 swizzle 会保留为 ``ComposeLayout``。Layout 中的 thread axis 表示元素归属于哪个执行成员,``TilePrimitiveCall.scope`` 则记录整项操作的协作层级。这些映射的具体含义与 lowering 过程见 :ref:`chap_tirx_tile_layout_lowering`。 + +这里需要区分两类 scope:``Buffer`` 数据指针上的 storage scope 描述数据位于 global、shared、local 或 TMEM;下一节的 execution scope 描述一项操作由哪一级线程集合执行。 + +执行层级保存在三个位置 +------------------------ + +执行层级在 IR 中分成区域边界、ID 声明和逐调用协作粒度: + +.. list-table:: + :header-rows: 1 + :widths: 28 32 40 + + * - 保存位置 + - 对应的 Python 写法 + - 节点中的信息 + * - device-entry ``AttrStmt`` + - ``Tx.device_entry()`` + - device region 的范围 + * - ``ScopeIdDefStmt.def`` + - ``tx = Tx.thread_id([128])`` + - ``def_ids``、``extents``、``scope`` 和可选的 ``preferred_extents`` + * - ``TilePrimitiveCall.scope`` + - ``Tx.tile.cta.copy(...)`` + - ``ExecScope.kind``,也就是这一条调用的协作层级 + +``ScopeIdDefStmt`` 记录当前区域可用的执行坐标及其范围,``ExecScope`` 记录当前调用的协作层级。同一组执行 ID 可以服务于多条具有不同 ``ExecScope`` 的 tile primitive,因此两部分信息各自保存。 + +``TilePrimitiveCall`` 保留 dispatch 所需信息 +--------------------------------------------- + +``TilePrimitiveCall`` 自身是一条 ``Stmt``,包含六个字段: + +.. list-table:: + :header-rows: 1 + :widths: 22 78 + + * - 字段 + - 保存的内容 + * - ``op`` + - tile operator 在 ``Op`` 注册表中的标识 + * - ``args`` + - ``Array``,可容纳 ``BufferRegion``、标量表达式和其他算子参数 + * - ``workspace`` + - 算子使用的预分配 buffer + * - ``config`` + - ``cta_group`` 等算子选项 + * - ``dispatch`` + - 可选的显式实现名称 + * - ``scope`` + - 保存协作粒度的 ``ExecScope`` + +Tile primitive 可能读取或改写多个 ``BufferRegion``,其结果和副作用通过这些区域传递,因此作为一条语句占据明确的执行位置。各操作按照自己的参数约定识别输入与输出;``TilePrimitiveCall`` 构造时还会检查 ``op`` 是否注册为 tile primitive,TIRx 的 visitor 也为它提供独立的分派入口。 + +``BufferRegion`` 的运行时类型保持为独立的引用对象;进入 ``PrimExpr`` 上下文时,可转换的定长区域会成为 ``BufferLoad``。``TilePrimitiveCall.args`` 使用 ``Array``,可以原样保留每一维的 ``Range(min, extent)``、底层 ``Buffer`` 引用及其 layout。 + +``tirx.Call`` 是 ``PrimExpr``,参数类型为 ``Array``,主要字段为 ``op``、``args`` 和 ``attrs``。``Tx.ptx.*`` 与 ``Tx.cuda.*`` 已经给出目标相关操作,parser 会用这种节点保存它们。返回标量的 ``Tx.ptx.elect_sync()`` 可以直接参与表达式;以副作用为主的调用则位于 ``Evaluate(Call)`` 中。 + +从打印结果验证这些关系 +------------------------ + +取得入门示例返回的 ``kernel`` 后,可以直接检查函数边界与语句树: + +.. code-block:: python + + print(type(kernel)) + print(kernel.params) + print(kernel.buffer_map) + print(type(kernel.body)) + print(next(iter(kernel.buffer_map.values())).layout) + + print(kernel.script( + syntax_sugar=False, + extra_config={"tirx.prefix": "Tx"}, + )) + +关闭 syntax sugar 后,printer 会展开参数 handle 与 buffer 绑定,并从 ``kernel.body`` 进入 visitor。默认 layout 在 script 输出中仍会省略,因此上面的 ``buffer.layout`` 可用于区分默认 layout 与空值。编写分析或变换 pass 时,可以使用 ``tvm.tirx.stmt_functor`` 中的 visitor 和 mutator。 + +``TilePrimitiveCall`` 将 tile 级语义保留为一条完整语句,``tirx.Call`` 则记录已经明确的调用及其返回类型。这些节点随后的转换顺序见 :ref:`chap_tirx_lowering_pipeline`,tile primitive 与 layout 的具体展开见 :ref:`chap_tirx_tile_layout_lowering`。 + +核心源码导航 +------------ + +- `function.h`_:``tirx.PrimFunc`` 的字段; +- `buffer.h`_:``Buffer``、layout 和 ``allocated_addr``; +- `buffer.py`_:``Buffer.view`` 如何构造共享 data pointer 的新视图; +- `layout.h`_:``Layout`` 层级、``Iter`` 与 ``Axis``; +- `stmt.h`_:平铺声明、``BufferRegion`` 与 ``ScopeIdDefStmt``; +- `tirx_stmt.h`_:``TilePrimitiveCall`` 的六个字段; +- `expr.h`_:``tirx.Call``、``tirx.Var`` 与 ``tirx.BufferLoad`` 等具体表达式节点; +- `ir_expr.h`_:通用 ``PrimExpr`` 基类与常量节点; +- `exec_scope.h`_:``ExecScope``、``ScopeIdDef`` 与 ``ScopeBinding``; +- `transform.cc`_:TIRx ``PrimFuncPass`` 的模块遍历边界; + +.. _function.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/function.h +.. _buffer.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/buffer.h +.. _buffer.py: https://github.com/apache/tvm/blob/v0.26.0/python/tvm/tirx/buffer.py +.. _layout.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/layout.h +.. _stmt.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/stmt.h +.. _tirx_stmt.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/tirx_stmt.h +.. _expr.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/expr.h +.. _ir_expr.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/ir/expr.h +.. _exec_scope.h: https://github.com/apache/tvm/blob/v0.26.0/include/tvm/tirx/exec_scope.h +.. _transform.cc: https://github.com/apache/tvm/blob/v0.26.0/src/tirx/ir/transform.cc diff --git a/zh/tirx_guide/arch/lowering_pipeline.rst b/zh/tirx_guide/arch/lowering_pipeline.rst index e542b031..4cc1405d 100644 --- a/zh/tirx_guide/arch/lowering_pipeline.rst +++ b/zh/tirx_guide/arch/lowering_pipeline.rst @@ -20,29 +20,7 @@ TIRx 编译流水线 =============== -调用 ``tvm.compile(mod, target, tir_pipeline="tirx")`` 时,输入并不是直接从 -Python 翻译成 PTX。编译器先消费 TIRx 特有的 tile primitive、execution scope -和 layout,再进行通用 TIR 的正规化、局部优化、类型合法化、host/device -拆分与代码生成。这里的 pass 是对 IR 执行一类转换、检查或标注的编译步骤; -target 指定设备和代码生成后端。 - -本页回答两个问题:**pipeline 里有哪些 passes,以及优化主要发生在哪里**。 -``Tx.gemm_async`` 如何使用 layout 推导 descriptor、直接 ``T.ptx.*`` 绕过什么, -单独放在 :ref:`chap_tirx_tile_layout_lowering` 中展开。 - -.. admonition:: 先给结论 - - 默认 TIRx 编译器可以概括为:**较薄的通用优化流水线,加上较厚的算子级 - lowering**。Kernel 作者或上层生成器负责 tile sizes、warp roles、同步、 - pipeline stages 和主要 layout;tile primitive 的实现负责合法性检查、 - descriptor/地址参数推导和指令分解;通用 passes 主要负责化简、合法化、 - 模块拆分和 ABI。 - -.. note:: - - 本书固定使用 Apache TVM ``0.26.0``。下面的 pass 名称和顺序都对应这个版本, - 并且讨论的是显式指定 ``tir_pipeline="tirx"`` 的路径。开发分支中的 pipeline - 可能已经变化;名为 ``"default"`` 的 pipeline 也不是 ``"tirx"`` 的别名。 +用 Python DSL 编写的 kernel 会先被解析为 TIRx ``PrimFunc``,其中保留 ``Tx.tile.*`` tile primitives、execution scopes 和 buffer layouts。调用 ``tvm.compile(mod, target, tir_pipeline="tirx")`` 后,``LowerTIRx`` 根据 target 选择 tile primitive 的具体实现,解析线程层级,并将 layout 转换为硬件指令参数或物理 buffer 地址。``tirx_pipeline`` 中余下的 passes 继续完成 IR 化简、类型合法化、host/device 拆分与 runtime ABI lowering,最终交给目标后端生成代码。 完整编译路径 ------------ @@ -56,14 +34,14 @@ target 指定设备和代码生成后端。 ▼ TIRx PrimFunc ├─ logical computation - ├─ Tx.* TilePrimitiveCall + ├─ Tx.tile.* TilePrimitiveCall ├─ Buffer + TileLayout └─ device / CTA / warpgroup / warp / thread scopes │ │ BindTarget ▼ LowerTIRx - ├─ TilePrimitiveDispatch:Tx.* → target-specific TIR / T.ptx.* + ├─ TilePrimitiveDispatch:Tx.tile.* → target-specific TIR / Tx.ptx.* └─ LowerTIRxCleanup:剩余 layout/access → physical buffer access │ ▼ @@ -77,23 +55,30 @@ target 指定设备和代码生成后端。 ▼ PTX / cubin / runtime module -``PrimFunc`` 是 TIR 中的函数表示,finalization 是代码生成前面向具体 target -的最后一组转换。``BindTarget`` 在 ``tirx_pipeline`` 之前运行。Target 不只是决定最后使用哪个 -code generator;``TilePrimitiveDispatch`` 在较早阶段就需要 target,才能查找 -对应的算子实现。``SplitHostDevice`` 位于 module-level pipeline 后半段,拆分后 -host 和 device functions 才分别进入各自的 finalization。 +``PrimFunc`` 是 TIR 中的函数表示,finalization 是代码生成前面向具体 target 的最后一组转换。``BindTarget`` 在 ``tirx_pipeline`` 之前运行。Target 同时服务于两个阶段:``TilePrimitiveDispatch`` 较早读取它来查找算子实现,最终 code generator 再读取它来生成目标代码。``SplitHostDevice`` 位于 module-level pipeline 后半段,拆分后 host 和 device functions 分别进入各自的 finalization。 高层 tile primitive 和直接 PTX 从不同位置进入这条路径: .. code-block:: text - Tx.gemm_async ── TilePrimitiveDispatch ──▶ T.ptx.tcgen05.mma ──┐ - ├─▶ 后续通用 passes - 直接 T.ptx.tcgen05.mma ────────────────────────────────────────┘ + 高层 tile primitive 路径 -因此,直接 PTX 绕过的是对应 tile primitive 的算子级选择、检查和参数推导, -不是整个 TIRx pipeline。两条路径的详细比较见 -:ref:`chap_tirx_tile_layout_lowering`。 + Tx.tile.gemm_async + │ TilePrimitiveDispatch + ▼ + Tx.ptx.tcgen05.mma + │ + ▼ + tirx_pipeline 后续 passes + + 直接 PTX 路径 + + Tx.ptx.tcgen05.mma + │ + ▼ + tirx_pipeline 后续 passes + +因此,直接 PTX 从对应 tile primitive 的算子级选择、检查和参数推导之后接入,并继续经过 ``tirx_pipeline`` 中余下的 passes。两条路径的详细比较见 :ref:`chap_tirx_tile_layout_lowering`。 ``LowerTIRx`` 的边界 ------------------------------ @@ -107,70 +92,55 @@ host 和 device functions 才分别进入各自的 finalization。 LowerTIRxCleanup, ]) -若设置 ``TVM_PRINT_AFTER_TIRX_DISPATCH_OPS``,两者之间还会临时插入一个 -``PrintIR``,但它只是调试 instrumentation,不是固定 transformation。 - -``TilePrimitiveDispatch`` 首先选择已注册的 target-specific 实现,把 -``Tx.copy``、``Tx.gemm_async``、``Tx.reduce`` 等 ``TilePrimitiveCall`` 替换为 -lower-level TIR 或 ``T.ptx.*``,并解析 device entry 内的 scope IDs。例如: +``TilePrimitiveDispatch`` 首先选择已注册的 target-specific 实现,把 ``Tx.tile.copy``、``Tx.tile.gemm_async``、``Tx.tile.reduce`` 等 ``TilePrimitiveCall`` 替换为 lower-level TIR 或 ``Tx.ptx.*``,并解析 device entry 内的 scope IDs。例如: .. code-block:: text - bx = T.cta_id([grid_x]) → bx = blockIdx.x - tx = T.thread_id([block_x]) → tx = threadIdx.x + bx = Tx.cta_id([grid_x]) → bx = blockIdx.x + tx = Tx.thread_id([block_x]) → tx = threadIdx.x -``LowerTIRxCleanup`` 随后对仍然存在的直接 ``BufferLoad`` / ``BufferStore`` 应用 -memory layout,展平相应 buffers,清除已消费的 layout metadata,并移除 -``tirx.buffer_offset`` wrappers。必须先 dispatch 后 cleanup:算子实现需要在 -metadata 消失前读取完整的 region、shape、dtype、scope 和 layout。 +``LowerTIRxCleanup`` 随后对仍然存在的直接 ``BufferLoad`` / ``BufferStore`` 应用 memory layout,展平相应 buffers,清除已消费的 layout metadata,并移除 ``tirx.buffer_offset`` wrappers。必须先 dispatch 后 cleanup:算子实现需要在 metadata 消失前读取完整的 region、shape、dtype、scope 和 layout。 ``LowerTIRx`` 成功结束后: - ``TilePrimitiveCall`` 已被具体实现替换; - 抽象 scope IDs 已变成 launch parameters、``Bind`` 和 thread bindings; - layout 已被算子 lowering 消费,或被 cleanup 物化为后端能处理的地址; -- ``T.ptx.*`` 等 target intrinsics 仍可存在; +- ``Tx.ptx.*`` 等 target intrinsics 仍可存在; - thread-binding loops 和 TIRx-specific loop annotations 仍可能存在,随后由 ``LowerTIRxOpaque`` 规范化; -- host/device split、ABI lowering 和最终 code generation 尚未完成。 - -因此,这个阶段只能说 **TIRx 的核心高层语义已经消解**,不能说已经得到最终 -PTX,也不宜把它描述成编译流程的终点。算子 dispatch 和 layout cleanup 的内部 -过程见 :ref:`chap_tirx_tile_layout_lowering`。 +- host/device split、ABI lowering 和最终 code generation 留给后续阶段。 最小示例:从 scope IDs 到 host/device --------------------------------------- -下面用一个不依赖 Tensor Core 的 scale kernel 串起这些阶段。它启动 4 个 CTAs, -每个 CTA 有 256 个 threads: +下面用一个执行逐元素计算的 scale kernel,跟踪 scope IDs 经过 ``LowerTIRx``、``SplitHostDevice`` 和 CUDA codegen 后的变化。这个 kernel 启动 4 个 CTAs,每个 CTA 有 256 个 threads: .. code-block:: python import tvm - from tvm.script import tirx as T - - @T.prim_func - def scale(A_ptr: T.handle, B_ptr: T.handle): - A = T.match_buffer(A_ptr, (1024,), "float32") - B = T.match_buffer(B_ptr, (1024,), "float32") - T.device_entry() - bx = T.cta_id([4]) - tx = T.thread_id([256]) + from tvm.script import tirx as Tx + + @Tx.prim_func + def scale(A_ptr: Tx.handle, B_ptr: Tx.handle): + A = Tx.match_buffer(A_ptr, (1024,), "float32") + B = Tx.match_buffer(B_ptr, (1024,), "float32") + Tx.device_entry() + bx = Tx.cta_id([4]) + tx = Tx.thread_id([256]) i = bx * 256 + tx - B[i] = A[i] * T.float32(2.0) + B[i] = A[i] * Tx.float32(2.0) -``T.device_entry()`` 标出 device region。``LowerTIRx`` 将抽象 IDs 解析成 -``blockIdx.x`` 和 ``threadIdx.x``;省略 buffer declarations 后,核心结构可写成: +``Tx.device_entry()`` 标出 device region。``LowerTIRx`` 将抽象 IDs 解析成 ``blockIdx.x`` 和 ``threadIdx.x``;省略 buffer declarations 后,核心结构可写成: .. code-block:: python - with T.launch_thread("blockIdx.x", 4) as bx: - tx = T.launch_thread("threadIdx.x", 256) + with Tx.launch_thread("blockIdx.x", 4) as bx: + tx = Tx.launch_thread("threadIdx.x", 256) i = bx * 256 + tx - B[i] = A[i] * T.float32(2.0) + B[i] = A[i] * Tx.float32(2.0) -这仍是 TIR 的结构摘要,不是 CUDA 源码,也不是完整的 printer output。 -``SplitHostDevice`` 随后把原来的一个 ``PrimFunc`` 拆成两部分: +这里展示的是省略 buffer declarations 后的 TIR 结构摘要。完整 printer output 还会包含声明等细节,CUDA 源码则由后续 codegen 生成。``SplitHostDevice`` 随后把原来的一个 ``PrimFunc`` 拆成两部分: .. code-block:: text @@ -180,8 +150,7 @@ PTX,也不宜把它描述成编译流程的终点。算子 dispatch 和 layout device scale_kernel └─ 每个 GPU thread 将一个元素乘以 2 -``MakePackedAPI`` 把 host entry 降低为 runtime 使用的统一 ABI(函数调用约定)。 -Device function 则交给 CUDA backend,最终生成与下面代码等价的 kernel: +``MakePackedAPI`` 把 host entry 降低为 runtime 使用的统一 ABI(函数调用约定)。Device function 则交给 CUDA backend,最终生成与下面代码等价的 kernel: .. code-block:: cuda @@ -193,9 +162,7 @@ Device function 则交给 CUDA backend,最终生成与下面代码等价的 ke 后续 passes:按职责理解 ------------------------ -``LowerTIRx`` 后的 passes 可以先分成四组;组内项目仍按源码定义的实际顺序 -执行,不能因为教学上分组就任意交换。``PassContext`` 是控制 pass 行为的 -编译配置对象。 +``LowerTIRx`` 后的 passes 可以先分成四组;组内项目保持源码定义的实际执行顺序。``PassContext`` 是控制 pass 行为的编译配置对象。 .. list-table:: :header-rows: 1 @@ -203,7 +170,7 @@ Device function 则交给 CUDA backend,最终生成与下面代码等价的 ke * - 职责 - 主要 passes - - 做什么、不做什么 + - 转换范围 * - 结构正规化 - ``UnifyThreadBinding``、``StmtSimplify``、``LowerTIRxOpaque``、 ``FlattenBuffer`` @@ -212,32 +179,24 @@ Device function 则交给 CUDA backend,最终生成与下面代码等价的 ke * - 局部程序变换 - ``NarrowDataType``、``VectorizeLoop``、``UnrollLoop``、第二次 ``StmtSimplify``、``CommonSubexprElim`` - - 缩窄安全的 index;兑现已有 vectorize/unroll 标记或配置;执行局部代数 - 与公共子表达式化简,而不是自动发现完整 schedule + - 缩窄安全的 index;兑现已有 vectorize/unroll 标记或配置;执行局部代数与公共子表达式化简;完整 schedule 由输入 IR 提供 * - 类型合法化 - BF16/FP8 compute legalization、BF16/FP8 storage legalization、 final ``LowerIntrin`` - - 将当前 target/backend 不能直接表示的 compute、storage 和 intrinsic - 改写为受支持形式;部分 pass 可能因 target 能力而成为 no-op + - 将当前 target/backend 需要降级表示的 compute、storage 和 intrinsic + 改写为受支持形式;具备原生表示能力的 target 保留原表示 * - 校验、模块和 ABI lowering - ``VerifyMemory``、``AnnotateEntryFunc``、``SplitHostDevice``、 ``LowerIket``、``MakePackedAPI`` - 验证设备计算处于 thread environment,抽取 device kernel,降低 kernel launch,并生成 runtime 可调用的 packed ABI -``VectorizeLoop`` 主要降低已经写成 ``T.vectorized`` 的 loops;它不会自己搜索 -应该向量化哪一层;禁用 vectorization 时,相应循环按标量循环处理。 -``UnrollLoop`` 默认主要处理显式 ``T.unroll``,也可以通过 -``PassContext`` 或 pragma 设置自动展开阈值。两者都不等于自动 tiling 或 -software-pipeline synthesis。 +``VectorizeLoop`` 主要降低已经写成 ``Tx.vectorized`` 的 loops,向量化层次来自输入 IR 中的标记;``PassContext`` 选择标量模式时,相应循环按标量形式处理。``UnrollLoop`` 默认主要处理显式 ``Tx.unroll``,也可以通过 ``PassContext`` 或 pragma 设置自动展开阈值。Tiling 和 software pipeline 的整体安排通常已经由 kernel 作者或上层生成器写入输入 IR。 Host/device split 与最终 codegen -------------------------------- -在 Apache TVM 0.26.0 中,``SplitHostDevice`` 是一个组合 pass:它识别 device -regions,将其抽取为 device ``PrimFunc``,并把 host 侧调用降低为 kernel-launch -约定。随后 ``MakePackedAPI`` 将公开的 host entry 改写为 TVM runtime 使用的 -packed-function ABI。 +在 Apache TVM 0.26.0 中,``SplitHostDevice`` 是一个组合 pass:它识别 device regions,将其抽取为 device ``PrimFunc``,并把 host 侧调用降低为 kernel-launch 约定。随后 ``MakePackedAPI`` 将公开的 host entry 改写为 TVM runtime 使用的 packed-function ABI。 .. code-block:: text @@ -248,7 +207,7 @@ packed-function ABI。 host launcher device PrimFunc ├─ 准备调用参数 ├─ thread_extent ├─ grid/block launch parameters ├─ physical buffer accesses - └─ 调用 device kernel └─ T.ptx.* / target intrinsics + └─ 调用 device kernel └─ Tx.ptx.* / target intrinsics │ │ │ host finalization │ device finalization ▼ ▼ @@ -265,15 +224,10 @@ Module-level pipeline 结束后,finalization 分别运行: - **host**:``LowerTVMBuiltin``、``LowerIntrin``; - **device**:``LowerWarpMemory``、``StmtSimplify``、``LowerIntrin``。 -``LowerTVMBuiltin`` 处理 ``tvm_*`` builtins,``LowerIntrin`` 处理 -target-specific intrinsics。 -``LowerWarpMemory`` 将 warp-scoped buffers 降成 local storage 和 shuffle 等 -形式。CUDA code generator 随后生成 CUDA C++;部分 ``T.ptx.*`` intrinsic 会被 -打印为 inline PTX 或 helper code。之后 NVRTC/NVCC/ptxas 等 CUDA toolchain -组件才继续产生可加载的 PTX 或 binary。``LowerIntrin`` 本身并不等于“输出 PTX”。 +``LowerTVMBuiltin`` 处理 ``tvm_*`` builtins,``LowerIntrin`` 处理 target-specific intrinsics。``LowerWarpMemory`` 将 warp-scoped buffers 降成 local storage 和 shuffle 等形式。CUDA code generator 随后生成 CUDA C++;部分 ``Tx.ptx.*`` intrinsic 会被打印为 inline PTX 或 helper code。之后 NVRTC/NVCC/ptxas 等 CUDA toolchain 组件继续产生可加载的 PTX 或 binary。``LowerIntrin`` 的产物仍由随后运行的 CUDA code generator 和 toolchain 继续处理。 -编译器优化什么、不优化什么 ----------------------------- +性能决策的职责边界 +------------------ 默认 TIRx pipeline 的责任边界如下: @@ -284,28 +238,21 @@ target-specific intrinsics。 * - 层次 - 主要责任 * - Kernel 作者或生成 kernel 的 agent - - 选择 tile sizes、pipeline stages、warp roles、execution scope、同步、 - 主要 layout 和跨 tile 的整体 schedule + - 选择 tile sizes、pipeline stages、warp roles、execution scope、同步、主要 layout 和跨 tile 的整体 schedule * - Layout helpers 与 tile dispatcher - - 根据显式 dtype/shape/mode 构造已知 layout;静态选择实现,检查合法性, - 推导 descriptor/物理参数,并将单个 tile operation 分解为硬件指令 - * - 通用 TIR passes + - 根据显式 dtype/shape/mode 构造已知 layout;静态选择实现,检查合法性,推导 descriptor/物理参数,并将单个 tile operation 分解为硬件指令 + * - ``LowerTIRx`` 后续 passes - 局部化简、index narrowing、显式或配置驱动的 vectorize/unroll,以及 dtype、module 和 ABI 合法化 * - CUDA backend 与 toolchain - 生成 target code,进行更底层的 peephole、register allocation 和组装 -所以,默认 pipeline 中确实有局部优化,但通常没有通用的自动 tiling、跨算子 -fusion、software-pipeline synthesis、warp-specialization synthesis、全局 -layout search、cost-model selection 或 autotuning。一个 dispatcher 根据 shape -将 tile operation 拆成多条硬件指令,是算子实现的一部分,不能与全程序 -schedule search 混为一谈。 +默认 pipeline 的自动优化集中在局部程序变换。Tile sizes、跨算子 fusion、software pipeline、warp specialization 和主要 layout 等整体 schedule 决策,通常由 kernel 作者或上层生成器提供。Dispatcher 则根据给定 shape 将一个 tile operation 映射成一组硬件指令;全程序 schedule search、cost-model selection 和 autotuning 属于更上层的生成与搜索系统。 检查完整流水线 -------------- -前面的 ``scale`` 例子也可以直接用于检查中间 IR。先绑定 CUDA device 与 LLVM -host target,再观察 ``LowerTIRx`` 前后的结果: +前面的 ``scale`` 例子也可以直接用于检查中间 IR。先绑定 CUDA device 与 LLVM host target,再观察 ``LowerTIRx`` 前后的结果: .. code-block:: python @@ -334,17 +281,12 @@ host target,再观察 ``LowerTIRx`` 前后的结果: ) print(exe.mod.imports[0].inspect_source()) -这个例子的 host module 只导入一个 device module,因此 ``imports[0]`` 就是 -CUDA module。如果问题发生在 ``Tx.*`` 到 ``T.ptx.*`` 之间,应进一步单独观察 -``TilePrimitiveDispatch``;Blackwell target 的设置和具体方法见 -:ref:`chap_tirx_tile_layout_lowering`。 +这个例子的 host module 恰有一个 device-module import,因此 ``imports[0]`` 对应 CUDA module。如果问题发生在 ``Tx.tile.*`` 到 ``Tx.ptx.*`` 之间,应进一步观察 ``TilePrimitiveDispatch``;Blackwell target 的设置和具体方法见 :ref:`chap_tirx_tile_layout_lowering`。 Pass 顺序参考 ------------- -Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时按下面 -顺序执行,共 19 步。若设置 ``tir.disable_cse_tir=True``,第 11 步会被省略, -后续 passes 依次前移。 +Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时按下面顺序执行,共 19 步。设置 ``tir.disable_cse_tir=True`` 时执行 18 步序列,后续 passes 的编号依次前移。 .. list-table:: :header-rows: 1 @@ -364,32 +306,33 @@ Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时 - 使用 arithmetic analyzer 化简 statements 和索引表达式 * - 4 - ``LowerTIRxOpaque`` - - 转换 thread-binding loops、消除无 annotation 的 unit loops,并规范化 - loop pragmas + - 转换 thread-binding loops、保留带 annotation 的 unit loops、折叠其余 + unit loops,并规范化 loop pragmas * - 5 - ``FlattenBuffer`` - 展平剩余 ``BufferLoad`` / ``BufferStore`` * - 6 - ``BF16ComputeLegalize`` - - target 不支持时将 BF16 compute 提升至 ``float32`` 并改写 + - 在需要 fallback 的 target 上将 BF16 compute 提升至 ``float32`` 并改写 * - 7 - ``NarrowDataType(32)`` - 能够证明安全时,将 index/loop expressions 缩窄到 32 bits * - 8 - ``VectorizeLoop`` - - Lower 已标记的 vectorized loops;禁用时将其作为标量 loops 处理 + - Lower 已标记的 vectorized loops;标量模式下按标量 loops 处理 * - 9 - ``UnrollLoop`` - - 展开显式 ``T.unroll``,并执行配置或 pragma 允许的自动展开 + - 展开显式 ``Tx.unroll``,并执行配置或 pragma 允许的自动展开 * - 10 - ``StmtSimplify`` - 在 vectorize/unroll 暴露常量后再次化简 * - 11 - ``CommonSubexprElim`` - - 执行可配置关闭的公共子表达式消除 + - 执行公共子表达式消除,是否运行由 ``tir.disable_cse_tir`` 控制 * - 12 - ``FP8ComputeLegalize`` - - target 不支持时将 FP8 compute 提升至默认的 ``float32`` 并改写 + - 在需要 fallback 的 target 上将 FP8 compute 提升至默认的 ``float32`` + 并改写 * - 13 - ``VerifyMemory`` - 检查 GPU target/default calling-convention function 中,参数 buffer 的 @@ -402,7 +345,7 @@ Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时 - 标注/抽取 device functions,并 lowering host-to-device kernel calls * - 16 - ``LowerIket`` - - IKET 未启用时移除相关 annotations;启用时生成所需 tracing 形式 + - 根据 IKET 开关移除相关 annotations,或生成所需 tracing 形式 * - 17 - ``MakePackedAPI`` - 将 host entry 改写为 runtime packed-function ABI @@ -413,8 +356,7 @@ Apache TVM 0.26.0 的 ``tirx_pipeline`` 在 **默认配置且 CSE 开启** 时 - ``BF16StorageLegalize`` - target 需要 fallback 时,将 BF16 storage 改写为等宽 ``uint16`` 表示 -不要仅凭 pass 名称推断默认会进行 aggressive auto-vectorization 或 -auto-unrolling;还需要检查 loop annotations 和当前 ``PassContext`` 配置。 +判断默认的 aggressive auto-vectorization 或 auto-unrolling 行为时,需要同时检查 pass 名称、loop annotations 和当前 ``PassContext`` 配置。 版本与核心源码 -------------- diff --git a/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst b/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst index 09543c94..36aaaf10 100644 --- a/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst +++ b/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst @@ -20,412 +20,318 @@ Tile Primitive 与 Layout Lowering ================================= -上一页 :ref:`chap_tirx_lowering_pipeline` 给出了完整 pipeline。本页只深入 -``LowerTIRx``,并沿一条 Blackwell ``Tx.gemm_async`` 解释 layout、算子静态分派 -和直接 PTX 之间的关系。 +上一页 :ref:`chap_tirx_lowering_pipeline` 给出了完整的编译流水线。本页进入 ``LowerTIRx`` 内部,沿一条 Blackwell ``Tx.tile.gemm_async`` 说明 tile primitive 和 layout 怎样逐步变成 target intrinsic 与物理地址。 -.. admonition:: 三个最容易混淆的结论 +``Tx.tile.gemm_async(C, A, B)`` 记录一次逻辑 tile GEMM。它给出操作数、区域和配置,具体使用哪些硬件指令、每个操作数怎样编码,则由 ``LowerTIRx`` 结合 target 与 layout 确定。这里的 layout 是一张从 **逻辑坐标** 到 **物理坐标** 的映射:生产者用它摆放数据,消费者用它找到同一份数据。 - 1. 默认 row-major layout 的地址可能与传统连续数组完全相同,但 **“映射结果 - 相同”不等于“IR 中没有 layout metadata”**。 - 2. ``Tx.gemm_async`` 不靠 layout 执行矩阵乘法;它读取 layout 以匹配/验证 - operand 的物理约定,并推导 descriptor、offset 和硬件指令参数。 - 3. 直接 ``T.ptx.*`` 省掉的是对应 ``Tx.*`` 的算子级 lowering。周围的 scope - lowering、普通地址展开、合法化、host/device split 和 codegen 仍然存在。 +``LowerTIRx`` 由两个顺序执行的 pass 组成: -贯穿示例:一条 ``Tx.gemm_async`` --------------------------------- +.. code-block:: text -下面抽取 :ref:`chap_tirx_primer` 中单-tile GEMM 的关键部分。完整 kernel 还 -包含 SMEM/TMEM allocation、barrier 初始化、等待和释放;这里仅保留与 lowering -有关的声明与 tile operations: + TilePrimitiveCall + Buffer layouts + target + │ + ▼ + TilePrimitiveDispatch + 选择实现,生成 Tx.ptx.* 与地址表达式 + │ + ▼ + LowerTIRxCleanup + 物化剩余普通访问,清除 layout metadata + │ + ▼ + 包含 Tx.ptx.* 与物理 Buffer 访问的 PrimFunc -.. code-block:: python +前一个 pass 读取算子级语义,后一个 pass 处理仍留在普通 ``BufferLoad``、``BufferStore`` 和指针访问中的 layout。同一个 buffer 的 layout 可以依次服务于这两个阶段。 - BLK_M, BLK_N, BLK_K = 128, 128, 64 +先读懂 Layout +------------- - A_layout = mma_shared_layout( - "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_M, BLK_K) - ) - B_layout = mma_shared_layout( - "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_N, BLK_K) - ) +一个逻辑 tile 可以落在普通内存、多个线程的局部存储或 TMEM 中。三类位置需要三种物理坐标: - Asmem = pool.alloc((BLK_M, BLK_K), "float16", layout=A_layout) - Bsmem = pool.alloc((BLK_N, BLK_K), "float16", layout=B_layout) +.. list-table:: + :header-rows: 1 + :widths: 22 35 43 - tmem = T.decl_buffer( - (128, 512), - "float32", - scope="tmem", - allocated_addr=tmem_addr[0], - layout=TileLayout(S[(128, 512) : (1@TLane, 1@TCol)]), - ) + * - 存储位置 + - Layout 映射 + - 物理坐标的含义 + * - GMEM / SMEM + - ``(i, j) → m`` + - ``m`` 是一个线性存储位置,也可以包含 padding 或 swizzle + * - 每线程局部存储 + - ``(i, j) → (thread axis, m)`` + - thread axis 指定持有元素的线程,``m`` 指定该线程内的局部位置 + * - Blackwell TMEM + - ``(i, j) → (TLane, TCol)`` + - 两个坐标共同指定 TMEM 中的硬件位置 - if warp_id == 0: - if T.ptx.elect_sync(): - Tx.gemm_async( - tmem[:, :BLK_N], Asmem[:, :], Bsmem[:, :], - accum=False, dispatch="tcgen05", cta_group=1, - ) +这些映射都由 ``TileLayout`` 表达。完整的 layout 代数、``S[...]`` 语法与组合规则见 :ref:`chap_tirx_layout_api`;这里关注它们在 lowering 中提供的信息。 - Dreg = T.alloc_local((BLK_N,), "float32") - Dreg_wg = Dreg.view( - 128, - BLK_N, - layout=TileLayout(S[(128, BLK_N) : (1@tid_in_wg, 1)]), - ) - Tx.wg.copy_async(Dreg_wg[:, :], tmem[:, :BLK_N]) +普通内存:逻辑下标映射到地址 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -这些语句已经给 lowering 提供了大部分 schedule: +省略 ``layout=`` 时,buffer API 默认创建连续行优先(row-major)layout。例如 shape 为 ``(4, 8)`` 的 buffer 使用: -.. list-table:: - :header-rows: 1 - :widths: 29 71 +.. code-block:: text - * - 输入信息 - - 它表达什么 - * - A/B/C 的 regions - - 本次 GEMM 的逻辑 ``M``、``N``、``K`` 范围 - * - A/B 的 SMEM layouts - - shared-memory 中的字节排列,以及选定的 128-byte swizzle - * - C 的 TMEM layout - - 声明期望的 accumulator datapath;dispatcher 可从特定 layout 推断 - ``.ws``,随后验证它与最终 tcgen05 datapath 一致并提取 slice offsets - * - ``warp_id`` 和 ``elect_sync`` - - 将 issuing scope 限定到一个被选中的 thread - * - ``dispatch="tcgen05"`` - - 强制选择 Blackwell tcgen05 variant - * - ``Dreg_wg`` layout - - 声明 readback 后每个逻辑元素的 thread ownership 和局部 slot - -其他 copy、reduce 等 tile primitive 也使用相同的 dispatcher 框架,但各自读取的 -layout 字段、约束和 lowering 结果并不相同。本例不能代表所有算子的具体硬件规则。 - -``TilePrimitiveDispatch`` 如何选择实现 ------------------------------------------------- - -前端把 ``Tx.gemm_async(...)`` 表示为 ``TilePrimitiveCall``。一次 dispatch 的输入 -由两部分合成: - -- ``TilePrimitiveCall`` 携带 operator、operand ``BufferRegion``、config 和 - ``dispatch=``; -- ``DispatchContext`` 提供 target、当前 execution scope、launch parameters、 - variable ranges 和插入初始化/分配语句所需的 callbacks。 - -候选实现按 ``(operator name, target kind)`` 注册,再按固定 priority 和 variant -名称排序。显式 ``dispatch="tcgen05"`` 只保留该 variant;没有显式指定时, -dispatcher 依次检查 predicates,并使用第一个成功返回 ``PrimFunc`` 的实现。 -这是 **静态规则选择**,不是 cost model,也不是 autotuning。 - -选中的 ``PrimFunc`` body 替换原 ``TilePrimitiveCall``。实现还可以通过 callbacks -请求 private allocation、device initialization、host initialization,或紧跟 -某个 buffer definition 的语句。因此,一次 lowering 不一定只在调用位置插入 -几条 PTX;它也可能准备 descriptor 或其他依赖资源。 - -这个 pass 还解析 ``T.device_entry()`` 内的抽象 scope IDs,生成 launch parameters、 -``Bind`` 和 ``thread_extent``。例如: + TileLayout(S[(4, 8)]) -.. code-block:: text + B[i, j] → B_flat[i * 8 + j] - T.cta_id([grid_x]) / T.thread_id([block_x]) - ↓ - blockIdx.x / threadIdx.x +这个公式与普通二维连续数组一致。layout metadata 会在编译前半程保留,因此 dispatcher 仍可对它执行 slice、匹配与检查。显式 ``layout=None`` 会让该字段保持为空,普通访问随后沿 shape/stride 规则展开。 -其中 ``grid_x`` 和 ``block_x`` 来自作者声明的 execution hierarchy,dispatcher -不会替 kernel 搜索 block size。 +物理排列也可以带 padding 或 swizzle。例如: -``Tx.gemm_async`` 的四步 lowering ---------------------------------------------- +.. code-block:: text + + TileLayout(S[(4, 8) : (16, 1)]) -第一步:slice layout 并验证 operands -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + B[i, j] → B_flat[i * 16 + j] -Dispatcher 从三个 ``BufferRegion`` 取得 extents,并对各 buffer layout 执行 -``slice`` 与 ``canonicalize``。随后检查: +这里每一行占 16 个元素,逻辑 shape 仍是 ``4×8``。对应的 backing allocation 需要覆盖 ``layout.span() = 56`` 个元素。``ComposeLayout`` 还可以把 XOR、shift 和 mask 组成 shared-memory swizzle。 -- C 是否位于 TMEM,A/B 是否位于该 variant 支持的 memory scope; -- operand dtype 和逻辑 ``M/N/K`` 是否满足指令约束; -- A/B 的 sliced layout 是否包含受支持的 SMEM atom、swizzle 和 alignment; -- C 的 sliced layout 是否匹配受支持的 TMEM datapath。 +每线程局部存储:逻辑下标映射到持有者和局部位置 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -这里不会把不兼容的 layout 自动“优化正确”。无法证明 layout、shape 或 -alignment 合法时,该 variant 会被拒绝;没有其他候选成功时,dispatch 失败。 +下面的 distributed layout 把第 ``m`` 行交给 workgroup 中的第 ``m`` 个线程,并把 ``n`` 作为该线程的局部位置: -第二步:A/B layout 变成 matrix descriptors -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +.. code-block:: python -``tcgen05`` 从 SMEM 读取 A/B。Dispatcher 将 sliced layout 与硬件支持的 K-major -和 MN-major swizzle atoms 匹配,从匹配结果取得: + Dreg_wg = Dreg.view( + 128, + N, + layout=TileLayout(S[(128, N) : (1@tid_in_wg, 1)]), + ) .. code-block:: text - swizzle mode - leading-dimension offset (ldo) - stride-dimension offset (sdo) - K-major / MN-major - 当前 MMA tile 相对 buffer 原点的 16-byte offset + Dreg_wg[m, n] → { tid_in_wg: m, local slot: n } -这些字段与 shared-memory base address 一起构成 matrix descriptor。前面的 -``Tx.cta.copy`` 或 ``Tx.copy_async`` 与后面的 MMA 因而通过同一个 layout contract -解释 SMEM 中的字节排列,kernel 作者不必手写 descriptor bit fields。 +``tid_in_wg`` 表达元素归属,默认 memory axis 记录该线程内的局部位置。CUDA toolchain 随后为这些局部值分配具体物理寄存器。 -第三步:验证 C datapath 并取得 TMEM 目标 -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +TMEM:逻辑下标映射到两个硬件坐标 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -示例 C layout 的正向映射是: +Blackwell TMEM 使用 ``TLane`` 和 ``TCol``: .. code-block:: text + TileLayout(S[(128, N) : (1@TLane, 1@TCol)]) + C[m, n] → { TLane: m, TCol: n } -它声明期望怎样解释 accumulator 的硬件坐标。Dispatcher 确定最终 tcgen05 -instruction datapath,验证 C layout 与它相容,并提取 sliced region 的 -``TLane`` / ``TCol`` offsets;若 region 从非零 column 开始,slice 会产生相应的 -``TCol`` offset。最后再与 ``allocated_addr`` 组合成目标 TMEM address。 +``TLane`` 表示 TMEM lane row,``TCol`` 表示 TMEM column。二者共同组成 tcgen05 指令使用的 TMEM 地址。这里的 ``TLane`` 属于存储坐标;CUDA ``lane_id`` 表示当前执行线程,两者含义不同。 -这里不能理解为“任意 C layout 都能改变硬件 accumulator 排列”。Layout-E -等受支持形式可以影响 ``.ws`` 推断,但最终仍必须匹配 tcgen05 能表达的 datapath。 +Layout 从哪里来 +~~~~~~~~~~~~~~~~ -第四步:shape、dtype 和 config 决定指令分解 -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +TIRx kernel 中的 layout 通常来自三处: -``Tx.gemm_async`` 表示整个逻辑 tile,不等于“一行 DSL 固定对应一条 PTX”。 -Dispatcher 主要根据逻辑 ``M/N/K``、dtype、``cta_group`` 和 operator config 选择 -合法 instruction shape;已有 layout 用来验证 operand 约定并取得每一步的 -descriptor offsets。 +1. **默认构造。** 省略 ``layout=`` 或写 ``layout="default"`` 时,parser 根据 shape 构造连续行优先 layout。 +2. **Helper 构造。** ``mma_shared_layout``、``tmem_datapath_layout`` 和 ``tcgen05_atom_layout`` 等 helper 根据 dtype、shape 与 mode 生成已知的硬件 layout。 +3. **显式声明。** 作者或上层生成器直接使用 ``TileLayout``、``ComposeLayout`` 等接口写出目标映射。 -对于示例中的 fp16 ``128×128×64`` GEMM,tcgen05 每个 K step 处理 16 个 fp16 -K elements,因此会产生 4 个 MMA iterations。去掉函数签名和大量常量细节后, -dispatch 后的结构可以概括为: +这里的自动化包含默认 layout 的构造,以及 lowering 根据既有 layout 推导指令参数。tile size、swizzle mode、线程分工和性能 layout 的选择通常由 kernel 作者、上层生成器或搜索系统完成。``LowerTIRx`` 接收这些已经确定的映射,再完成局部推导与合法性检查。 -.. code-block:: text +沿一条 ``Tx.tile.gemm_async`` 看 lowering +------------------------------------------- - # lowering sketch:尖括号内容不是可调用的 Python API - desc_a = <由 Asmem base、ldo、sdo、swizzle 组成的 TIR expression> - desc_b = <由 Bsmem base、ldo、sdo、swizzle 组成的 TIR expression> - desc_i = T.uint32() +下面抽取 :ref:`chap_tirx_primer` 中单-tile GEMM 的核心语句。SMEM/TMEM allocation、barrier 和等待代码在这里省略: - for ki in T.unroll(4): - T.ptx.tcgen05.mma( - , - , - , - desc_i, - enable_input_d=(ki != 0), - ... - ) +.. code-block:: python -``desc_i`` 对 dense tcgen05 路径是 dispatcher 在编译期间算出的 ``uint32`` -常量,不是 runtime 再调用某个 descriptor encoder。``T.unroll(4)`` 则由后面的 -``UnrollLoop`` 展开,随后的 ``StmtSimplify`` 会化简各次迭代的常量。由此可以看出: -instruction decomposition 来自 tcgen05 operator lowering,而 loop 展开属于 -后续通用 pass。 + BLK_M, BLK_N, BLK_K = 128, 128, 64 -Readback 延续同一个物理约定 -~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + A_layout = mma_shared_layout( + "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_M, BLK_K) + ) + B_layout = mma_shared_layout( + "float16", SwizzleMode.SWIZZLE_128B_ATOM, (BLK_N, BLK_K) + ) -tcgen05 按最终确定的 datapath 把结果写入 TMEM;C layout 用于验证并解释这个 -物理结果。随后的 ``Tx.wg.copy_async`` 同时读取 C layout 和 ``Dreg_wg`` layout, -选择匹配的 ``tcgen05.ld`` form,并把逻辑 ``(m, n)`` 分配给 -``tid_in_wg=m`` 的 thread 及其局部 slot ``n``: + Asmem = pool.alloc((BLK_M, BLK_K), "float16", layout=A_layout) + Bsmem = pool.alloc((BLK_N, BLK_K), "float16", layout=B_layout) -.. code-block:: text + C = Tx.decl_buffer( + (128, 512), + "float32", + scope="tmem", + allocated_addr=tmem_addr[0], + layout=TileLayout(S[(128, 512) : (1@TLane, 1@TCol)]), + ) - GMEM - │ Tx.cta.copy / Tx.copy_async:按 A/B layout 写入 - ▼ - swizzled SMEM - │ Tx.gemm_async:descriptor 按同一 layout 读取 - ▼ - TMEM (TLane, TCol) - │ Tx.wg.copy_async:按 C 与 register layouts 解释和读取 - ▼ - per-thread local slots (tid_in_wg, m) + if warp_id == 0: + if Tx.ptx.elect_sync(): + Tx.tile.gemm_async( + C[:, :BLK_N], + Asmem[:, :], + Bsmem[:, :], + accum=False, + dispatch="tcgen05", + cta_group=1, + ) -Layout 没有执行 copy 或 MMA;它是 producer 与 consumer 对“同一个逻辑元素 -位于哪里”的共同约定。 +这一条调用向 dispatcher 提供了四组信息: -Layout 的完整生命周期 ----------------------- +.. list-table:: + :header-rows: 1 + :widths: 30 70 -同一个 layout 从创建到消失会经过下面几个阶段: + * - 输入 + - Lowering 中的用途 + * - C/A/B 的 ``BufferRegion`` + - 确定本次调用的逻辑 ``M``、``N``、``K`` 范围及各 region 的起点 + * - A/B 的 SMEM layouts + - 确定 shared-memory 排列、swizzle、matrix descriptor 参数和每个指令 tile 的偏移 + * - C 的 TMEM layout + - 验证 accumulator datapath,并取得 region 的 ``TLane`` 与 ``TCol`` 偏移 + * - dtype、``cta_group`` 与 ``dispatch`` + - 约束候选实现、指令 shape 和 tcgen05 variant -.. code-block:: text +选择实现 +~~~~~~~~ - parser 默认构造 / helper synthesis / 作者显式构造 - │ - ▼ - attach 到 Buffer - │ - ▼ - Buffer view / region slice - │ - ▼ - canonicalize / match - │ - ┌───────────┴───────────┐ - ▼ ▼ - operator dispatcher LowerTIRxCleanup - 消费硬件语义与 offsets 物化剩余的 memory offset - └───────────┬───────────┘ - ▼ - layout metadata 被清除 - │ - ▼ - backend 只看到地址和 intrinsics +Python 前端把 ``Tx.tile.gemm_async`` 保存为 ``TilePrimitiveCall``。候选实现以 ``(operator, target kind)`` 为键注册,并带有 variant 名称、优先级和适用条件。 -Dispatcher 和 cleanup 不是互斥的二选一。同一个 shared-memory swizzle 可以被 -GEMM dispatcher 用来构造 descriptor,同时 cleanup 仍会把这个 buffer 上残留的 -普通 ``BufferLoad`` / ``BufferStore`` 展开成物理地址。 +显式 ``dispatch="tcgen05"`` 会筛选出对应 variant;省略 ``dispatch=`` 时,dispatcher 按优先级检查各候选的适用条件。选中的实现返回一个 ``PrimFunc``,其函数体替换原来的 ``TilePrimitiveCall``。整个过程发生在编译期,依据 target、操作数 region、layout、dtype 和 config 做规则匹配。 -.. list-table:: - :header-rows: 1 - :widths: 26 37 37 - - * - Layout 类型 - - Operator dispatcher 怎样使用 - - Cleanup 怎样使用 - * - 单一 memory axis ``m``,可含 swizzle - - 若 buffer 参与 tile primitive,可读取它以推导 descriptor、vector width - 或 offsets - - 将剩余直接 access 物化为一个线性 offset - * - ``laneid`` / ``tid_in_wg`` 加局部 ``m`` - - Register-aware operator 将它解释为 thread ownership 与局部 slot - - 不能直接遗留;必须先通过理解它的 operator,或用 ``.view()`` 后再以 - ``.local()`` 取得当前 thread 的 storage view - * - ``TLane`` / ``TCol`` - - TMEM-aware operator 验证 datapath 并取得硬件地址 - - 不能被普通 TIR ``BufferLoad`` 压成一个 pointer offset - -默认 row-major 不等于“没有 layout” -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ - -在本章讨论的 CUDA TIRx ``PrimFunc`` 中,``T.match_buffer``、``T.decl_buffer`` -等 buffer APIs 的 ``layout`` 参数默认值就是 ``"default"``。因此作者省略 -``layout=`` 时,parser 会自动构造: +其他 copy、reduce 和同步类 tile primitive 也经过同一套选择流程。每个实现自行定义要读取哪些 layout 坐标、接受哪些 shape,以及最终生成哪些 target intrinsics。 -.. code-block:: text +切出本次调用的 Layout +~~~~~~~~~~~~~~~~~~~~~~ - TileLayout(S[shape]) +调用参数是 ``C[:, :BLK_N]``、``Asmem[:, :]`` 和 ``Bsmem[:, :]`` 这样的 ``BufferRegion``。Dispatcher 先对各 buffer layout 执行 ``slice``,把 region 起点折入 layout offset,再通过 ``canonicalize`` 得到便于匹配的等价形式。 -它是 dense row-major layout。三种写法的差别如下: +tcgen05 实现随后检查: -.. list-table:: - :header-rows: 1 - :widths: 24 35 41 - - * - 写法 - - 普通连续 access 的结果 - - IR 中保留的信息 - * - 省略 ``layout=``,或写 ``layout="default"`` - - cleanup 后与传统 row-major 地址相同 - - 有一个可供 slice、match 和 operator 检查的 ``TileLayout`` - * - ``layout=None`` - - 通过普通 shape/stride 规则也可得到相同地址 - - 明确不附带 layout metadata;operator 无法从它取得专用映射 - * - 显式硬件 layout - - 地址可能包含 padding/swizzle,或映射到 thread/TMEM axes - - 携带 operator-specific 的物理约定 - -所以,对普通连续 buffer 来说,“默认 layout”和“没有 layout”可能产生完全 -相同的最终地址;差别在于编译器前半程有没有一份统一、可检查的映射契约。 -它本身不会凭空带来性能收益,更不等于编译器已经选出了最佳 layout。 - -Layout 自动到什么程度 -~~~~~~~~~~~~~~~~~~~~~~ +- C 的 scope、dtype、shape 与 ``TLane/TCol`` 映射; +- A/B 的 scope、dtype、alignment 与 shared-memory atom; +- A/B 的 major mode 和 swizzle 是否落在该指令支持的组合中; +- ``M/N/K`` 与 ``cta_group`` 是否可以分解成合法的 tcgen05 指令 shape。 -“自动 layout”常混用三种含义: +这些检查把 layout 当作调用约束。layout 与候选实现的硬件约定一致后,lowering 才继续生成 descriptor 和指令。 -1. **默认构造。** Parser 在省略参数时补 dense row-major layout。这只是默认 - 语义,不是性能搜索。 -2. **Helper synthesis。** ``mma_shared_layout``、``tmem_datapath_layout``、 - ``tcgen05_atom_layout`` 等 helper 根据显式 dtype、shape 和 mode 构造已知 - 硬件 layout。选择哪个 helper 和 mode 仍由作者或上层生成器决定。 -3. **Lowering-time parameter inference。** Dispatcher 结合既有 layout 与 - shape、dtype、``cta_group`` 和 config,推导 major mode、descriptor fields、 - instruction shape 与 offsets。它是从已给定约定推出底层参数,不是反向搜索 - 最优 layout。 +A/B Layout 推导 matrix descriptor +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -默认 TIRx pipeline 没有一个全局 ``InferOptimalLayout`` pass,也不会从任意 -手写 PTX 反推出 lane/register、SMEM swizzle 或 TMEM layout。 +tcgen05 从 SMEM 读取 A/B 时,硬件通过 matrix descriptor 解释 shared-memory 地址。Dispatcher 将 sliced layout 与受支持的 K-major 或 MN-major swizzle atom 匹配,并得到: -三种 layout 怎样变成物理位置 ------------------------------ +.. code-block:: text -普通 memory layout:逻辑下标变成地址 -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + major mode + swizzle mode + leading-dimension offset (ldo) + stride-dimension offset (sdo) + 当前指令 tile 相对 region 原点的 16-byte offset -考虑逻辑 shape ``(4, 8)``、row stride 为 16 的 padded layout: +``ldo``、``sdo`` 和 swizzle 等字段可在 dispatch 时确定;SMEM base address 属于运行时值。生成的 TIR 会调用 ``Tx.ptx.tcgen05.encode_matrix_descriptor``,把运行时地址与这些字段编码到一个局部 ``uint64`` descriptor 中。后续各个 MMA iteration 再加入以 16 bytes 为单位的 tile offset。 -.. code-block:: text +这也解释了 producer 与 consumer 为何要共享 layout。前面的 copy 按 A/B layout 把元素写入 swizzled SMEM,GEMM dispatcher 从同一份 layout 生成 descriptor,tcgen05 因而按照相同的字节排列读回元素。 - layout = TileLayout(S[(4, 8) : (16, 1)]) +C Layout 推导 TMEM 地址 +~~~~~~~~~~~~~~~~~~~~~~~ - layout.apply(i, j)["m"] = i * 16 + j +示例中 C region 的映射为: - B[i, j] → B_flat[i * 16 + j] +.. code-block:: text -这是纯映射示例;实际 backing allocation 必须至少容纳 -``layout.span() = 56`` 个 elements,而不是只分配逻辑元素数 ``4 * 8 = 32``。 -若通过函数参数传入 B,caller 也必须满足这个物理容量约定。 + C[m, n] → { TLane: m, TCol: n } -``LowerTIRxCleanup`` 会先把 layout 产生的物理坐标写入 access index,并保留 -buffer 的 ``elem_offset`` metadata;后续 ``FlattenBuffer`` 再把 -``elem_offset`` 折入最终线性 index。所以上图表示的是最终 **有效地址语义**, -不是声称 cleanup 内某一个 AST 节点已经完成所有后续 folding。 +对 layout 做 slice 后,dispatcher 分别取得 ``TLane`` offset 和 ``TCol`` offset。它们与 ``allocated_addr`` 一起传给 ``Tx.cuda.get_tmem_addr``,形成 tcgen05 使用的 TMEM 目标地址: -如果使用 ``ComposeLayout``,物理 offset 还可能包含 shared-memory swizzle 的 -XOR、shift 和 mask。完整 layout 代数见 :ref:`chap_tirx_layout_api`。 +.. code-block:: text -Distributed layout:逻辑元素变成 ownership 与局部 slot -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + get_tmem_addr(allocated_addr, TLane offset, TCol offset) -Register-backed tile 通常需要两个物理坐标:哪个 thread 持有元素,以及它在该 -thread 的第几个局部 slot。例如: +当 C region 从非零行或非零列开始时,对应偏移会进入这两个坐标。Dispatcher 同时验证 sliced layout 与所选 accumulator datapath 一致。 -.. code-block:: python +Shape 与 dtype 决定指令分解 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +示例表示一次 fp16 ``128×128×64`` GEMM。tcgen05 的这个 dense 路径每个 K step 处理 16 个 fp16 元素,因此 K 方向被分成 4 次 MMA。省略具体 TIR 语法后,dispatch 结果的结构如下: + +.. code-block:: text - fragment_layout = TileLayout( - S[(8, 4, 2) : (4@laneid, 1@laneid, 1)] + # 结构示意;尖括号中的内容代表生成的 TIR expression + desc_a = encode_matrix_descriptor( + , , , ) + desc_b = encode_matrix_descriptor( + , , , + ) + desc_i = Tx.uint32() + + for ki in Tx.unroll(4): + Tx.ptx.tcgen05.mma( + get_tmem_addr(, , ), + add_16B_offset(desc_a, ), + add_16B_offset(desc_b, ), + desc_i, + enable_input_d=(ki != 0), + ... + ) -把它解释成逻辑 ``8×8`` tile 时: +这里的 ``desc_a`` 和 ``desc_b`` 需要运行时 SMEM 地址,所以 descriptor encoding 保留在生成的 TIR 中。dense ``desc_i`` 只依赖指令 shape、dtype、major mode 与 ``cta_group``,dispatcher 会把它折叠为编译期 ``uint32`` 常量。 -.. code-block:: text +``Tx.unroll(4)`` 会由后续 ``UnrollLoop`` 展开,``StmtSimplify`` 再化简每次迭代中的常量表达式。指令数量与操作数分块来自 tile primitive 实现,循环展开和局部化简由流水线后段完成。 - laneid = 4 * row + col // 2 - m = col % 2 +Readback 使用另一组 Layout +~~~~~~~~~~~~~~~~~~~~~~~~~~ -``laneid`` 表示 ownership,``m`` 表示 lane-local slot;``m`` 不是最终 PTX 中 -某个固定寄存器编号。真实 register allocation 仍由 CUDA toolchain 完成。 +GEMM 将结果写入 TMEM 后,readback 是另一条独立的 tile primitive: -普通 TIR ``BufferLoad`` 无法只凭一个 offset 验证“当前 thread 是否拥有这个 -逻辑元素”。因此,含 thread axis 的 layout 如果直接遗留到 cleanup 会报错。 -``.view()`` 先建立 distributed logical view,随后还必须通过 ``.local()`` 取得 -当前 thread 的 storage view;``.view()`` 单独使用并不能让直接 load 合法。 +.. code-block:: python -TMEM layout:逻辑元素变成 ``TLane`` 与 ``TCol`` -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Dreg = Tx.alloc_local((BLK_N,), "float32") + Dreg_wg = Dreg.view( + 128, + BLK_N, + layout=TileLayout(S[(128, BLK_N) : (1@tid_in_wg, 1)]), + ) + Tx.tile.wg.copy_async(Dreg_wg[:, :], C[:, :BLK_N]) -Blackwell TMEM 使用二维硬件坐标: +它同时读取 C 的 ``TLane/TCol`` layout 和 ``Dreg_wg`` 的 distributed layout,选择相应的 ``tcgen05.ld`` form,并把 TMEM 元素送到指定线程的局部位置。整个数据流为: .. code-block:: text - TileLayout(S[(128, N) : (1@TLane, 1@TCol)]) + GMEM + │ copy primitive 按 A/B layout 写入 + ▼ + swizzled SMEM + │ gemm_async 按同一 layout 生成 descriptor 并读取 + ▼ + TMEM (TLane, TCol) + │ wg.copy_async 同时解释 TMEM 与 distributed layouts + ▼ + per-thread local slots (tid_in_wg, m) - C[m, n] → { TLane: m, TCol: n } +Layout 在这里充当 producer 与 consumer 之间的物理映射契约;tile primitive 负责把这份契约翻译成具体硬件操作。 -``TLane`` 是 TMEM 的物理 lane row,不是执行当前代码的 CUDA ``lane_id``。 -``TLane`` 与 ``TCol`` 也不是普通 pointer 的两个 strides;tcgen05-aware operator -必须先验证并解释它们。若这种二维坐标直接留给普通 ``BufferLoad``,cleanup -无法把它降成所要求的单一 memory offset。 +``LowerTIRxCleanup`` 处理剩余 Layout +------------------------------------ -``LowerTIRxCleanup`` 的准确边界 ---------------------------------------------- +``TilePrimitiveDispatch`` 已将全部 ``TilePrimitiveCall`` 展开。普通 ``BufferLoad``、``BufferStore`` 和指针访问仍可能引用带 layout 的 buffer,``LowerTIRxCleanup`` 随后处理这些访问: -Dispatcher 完成后,``LowerTIRxCleanup`` 运行 ``LayoutApplier``。对剩余普通 -``BufferLoad`` / ``BufferStore``,其核心工作是: +.. code-block:: text + + 创建或附着 Layout + │ + ▼ + Buffer view / region slice + │ + ▼ + TilePrimitiveDispatch 读取并展开算子 + │ + ▼ + LowerTIRxCleanup 映射剩余普通访问 + │ + ▼ + 重建物理 Buffer,清除 layout metadata + │ + ▼ + 后续 TIR passes 与 target codegen + +对普通 memory layout,``LayoutApplier`` 的核心过程可以写成: .. code-block:: text @@ -435,100 +341,87 @@ Dispatcher 完成后,``LowerTIRxCleanup`` 运行 ``LayoutApplier``。对剩余 layout.canonicalize().apply(indices, shape) │ ▼ - 一个 symbolic physical coordinate + 一个 symbolic memory coordinate │ ▼ - flattened buffer access - -对于 CUDA 的直接 memory access,有 layout 时最终必须只产生一个 physical -coordinate,通常命名为 ``m``;没有 layout 时则使用普通 shape/stride 规则。 -``LayoutApplier`` 还把 layout-backed buffers 重建为 physical views并清空 layout -metadata。随后 ``BufferOffsetRemover`` 消除 ``tirx.buffer_offset(BufferLoad)`` -wrapper,使其中的 offset 与已经展平的访问一致。 + physical buffer access -这里的 ``symbolic`` 表示编译器在 lowering 时构造、化简地址表达式;表达式中仍可 -包含 runtime loop index、``threadIdx`` 或函数参数,并非所有地址都在编译期变成 -常量。 +这里的 symbolic coordinate 可以包含运行时循环变量、``threadIdx`` 或函数参数。后续 ``FlattenBuffer`` 会把 buffer 的 ``elem_offset`` 合入线性 index,形成最终地址语义。``BufferOffsetRemover`` 则消除 ``tirx.buffer_offset(BufferLoad)`` wrapper,使其中的 offset 与物理访问保持一致。 -完整 ``LowerTIRx`` 成功后,可以依赖: +普通指针访问需要 layout 最终产生一个 memory coordinate,通常命名为 ``m``。带有 ``laneid``、``tid_in_wg``、``TLane`` 或 ``TCol`` 的映射表达额外的硬件坐标;它们应先由理解这些坐标的 tile primitive 消费,或通过 ``.view()`` 与 ``.local()`` 转成当前线程的存储视图。若这类坐标以普通 ``BufferLoad`` 留到 cleanup,编译器会给出诊断。 -- 所有 ``TilePrimitiveCall`` 已被具体实现替换,否则 pass 会失败; -- scope IDs 已被解析; -- layout metadata 已被 operator 消费或由 cleanup 物化并清除; -- ``T.ptx.*`` 可以继续存在,尚未变成最终 CUDA source/PTX assembly; -- 类型合法化、host/device split 和 ABI lowering 仍未完成。 +``LowerTIRx`` 完成后,所有 ``TilePrimitiveCall`` 都已展开,layout metadata 已被 dispatcher 读取或由 cleanup 物化,``Tx.ptx.*`` 等 target intrinsics 则继续进入后续 TIR passes 与 CUDA codegen。 -高层 tile primitive 与直接 PTX +高层 Tile Primitive 与直接 PTX ------------------------------- -``T.ptx.tcgen05.mma`` 是 target intrinsic ``Call``,不是 ``TilePrimitiveCall``。 -它不会进入 ``gemm_async`` 的 registered variants,但 surrounding kernel -仍会经过 scope lowering、cleanup、通用 passes、host/device split 和 codegen。 +直接写 ``Tx.ptx.tcgen05.mma`` 时,这个 target intrinsic ``Call`` 从一开始就在输入 IR 中。``TilePrimitiveDispatch`` 的算子分派针对 ``TilePrimitiveCall``,作者写下的 PTX intrinsic 会原样保留。高层写法生成的 PTX intrinsic 与直接写入的 PTX intrinsic,随后都经过 cleanup 和流水线后段: + +.. code-block:: text + + Tx.tile.gemm_async + │ + ▼ + TilePrimitiveDispatch + │ + ▼ + Tx.ptx.tcgen05.mma + │ + ▼ + cleanup → 后续 TIR passes → CUDA codegen + + 直接 Tx.ptx.tcgen05.mma + │ + ▼ + TilePrimitiveDispatch 保留该 intrinsic + │ + ▼ + cleanup → 后续 TIR passes → CUDA codegen + +两种写法的差别集中在 PTX intrinsic 生成之前: .. list-table:: :header-rows: 1 - :widths: 30 35 35 - - * - 责任 - - ``Tx.gemm_async`` - - 直接 ``T.ptx.tcgen05.mma`` - * - Backend variant - - Dispatcher 选择,或验证显式 ``dispatch=`` - - 作者已经选择 - * - Logical shape 到 instruction tiles - - Dispatcher 根据 shape/dtype/config 分解 - - 作者手写每条 instruction - * - Operand layout 合法性 - - Dispatcher 做 operator-specific matching 和检查 - - 主要由作者保证 - * - SMEM/TMEM descriptor 与 offsets + :widths: 27 37 36 + + * - 编程责任 + - ``Tx.tile.gemm_async`` + - 直接 ``Tx.ptx.tcgen05.mma`` + * - 实现选择 + - Dispatcher 根据 target 选择或验证 ``dispatch=`` + - 作者已经选定具体 intrinsic + * - 逻辑 tile 分解 + - 根据 shape、dtype 和 config 生成指令 tiles + - 作者逐条写出指令 + * - 操作数 layout 检查 + - Dispatcher 执行算子专用的匹配与验证 + - 作者按照 PTX 约定组织操作数 + * - Descriptor 与 TMEM 地址 - 从 layout、region 和 config 推导 - - 作者手写或调用低级 encoding intrinsics - * - Lane/local-slot/TMEM operand 顺序 - - Layout 与 dispatcher 共同表达和检查 - - 作者编码在 lane 公式、地址与 operand 顺序中 - * - Simplify、legalize、split、codegen - - 仍然执行 - - 仍然执行 - -“作者编码寄存器顺序”不是说作者指定 ``%r17`` 这样的最终物理寄存器号;作者 -只是在 local arrays、lane formulas 和 PTX operands 中表达值的相对位置,真正的 -register allocation 仍由 ptxas 等工具完成。区别在于:直接 PTX 的 IR 不再保留 -“这个 operand 对应逻辑 ``C[m,n]``”的完整 tile-op contract,dispatcher 因而无法 -替作者做同等级的结构匹配与 layout mismatch 诊断。 - -直接 PTX 也不必然绕过所有 memory layout: - -- 若 intrinsic 参数经 ``buf.ptr_to(logical_indices)`` 构造,且 buffer layout - 能降低为 **单一 memory-axis offset**,cleanup 仍会把逻辑 indices 映射成 - 物理地址。 -- Distributed/thread-axis layout 或 ``TLane/TCol`` layout 不能把 ``ptr_to`` - 当作通用逃生口;若它们仍以直接 ``BufferLoad`` 形式出现,cleanup 会报错。 -- 一旦 intrinsic 只接收 raw base address 和作者已经算好的 offset,layout - mapping 就已被绕过。``buf.data + raw_offset`` 是常见写法,但不是唯一方式。 - -因此,“直接 PTX 没有消灭 layout”的准确含义是:它只会让 **那个低级 -instruction call** 不再经过 tile-op layout inference;周围 buffers 的普通地址 -访问仍可能需要 layout。反过来,如果所有相关地址、lane mapping 和 operand -顺序都由作者以 raw expressions 写完,那么 layout 对这条指令当然不会再提供 -额外推导。 - -失败表示检查,而不是自动修复 ------------------------------- - -假设 C 位于 TMEM,却只给它一个普通 row-major ``m`` layout,然后强制 -``dispatch="tcgen05"``。该 layout 无法证明逻辑 C region 与 ``TLane/TCol`` -datapath 相容,tcgen05 implementation 会拒绝它;如果没有其他候选,编译器 -报告 dispatch failure。 - -这类错误说明了 TIRx 的责任边界:dispatcher 会从 **已经声明的** layout 推导 -底层参数并验证约束,但不会搜索一个新 layout,再悄悄重写 kernel 的 producer、 -consumer 和 allocation 使它们全部匹配。 - -检查 dispatch 与 layout lowering + - 作者调用低级编码与地址接口 + * - Lane 与局部值顺序 + - Distributed layout 表达归属和局部位置 + - 作者用 lane 公式、局部数组和 intrinsic 参数表达 + +直接 PTX 仍可与 layout 共存,具体取决于地址怎样构造: + +- PTX 参数来自 layout-backed memory buffer 的 ``ptr_to(logical_indices)`` 时,cleanup 可以把单一 memory-axis layout 映射成物理地址。 +- 使用 ``Tx.ptr_byte_offset`` 或 ``Tx.handle_add_byte_offset`` 提供 byte offset 时,地址表达式直接携带作者选定的物理映射。 +- TMEM 地址可以通过 ``Tx.cuda.get_tmem_addr(base, tlane, tcol)`` 明确给出,distributed 数据的 lane 归属与局部顺序也由作者明确组织。 + +因此,直接 PTX 省去了 tile primitive 提供的算子级推导与检查;layout 仍可负责周围 buffer 的地址映射。作者显式写出的 descriptor、TMEM 坐标、lane 公式和局部顺序承担同一组物理约定。ptxas 等工具继续负责 ``%r17`` 一类物理寄存器的最终分配。 + +Layout 约束怎样产生诊断 +----------------------- + +假设 C 位于 TMEM,却声明为只有 ``m`` 轴的连续行优先 layout,同时强制 ``dispatch="tcgen05"``。tcgen05 variant 期望 C region 映射到受支持的 ``TLane/TCol`` datapath;两组物理坐标发生冲突,该候选实现会被拒绝。 + +这类诊断来自 layout 的契约作用。Dispatcher 从已声明的 layout 推导底层参数并验证硬件约束,kernel 作者或上层生成器负责让 producer、consumer 与 allocation 使用一致的映射。 + +检查 Dispatch 与 Layout Lowering -------------------------------- -调试时可以在 cleanup 删除 layout metadata 前,单独查看 dispatch 结果: +需要定位问题时,可以分别打印 authored TIRx、dispatch 结果和完整 ``LowerTIRx`` 结果: .. code-block:: python @@ -550,39 +443,24 @@ consumer 和 allocation 使它们全部匹配。 lowered = TT.LowerTIRx()(bound) print(lowered.script()) -``TilePrimitiveDispatch`` 运行时必须能解析出 target;显式 ``BindTarget`` 最清晰, -也可以依赖已有的 PrimFunc target attribute 或 current target context。这里指定 -``sm_100a`` 是因为示例使用 Blackwell ``tcgen05``;生成和运行最终代码还要求相应 -CUDA toolkit 与硬件支持。 - -``LowerTIRx`` 也提供一个调试开关,在 dispatch 与 cleanup 之间打印 IR: - -.. code-block:: bash - - TVM_PRINT_AFTER_TIRX_DISPATCH_OPS=1 python your_kernel.py - -检查 ``Tx.gemm_async`` 时,建议依次确认: +示例使用 Blackwell ``tcgen05``,因此 target 设为 ``sm_100a``。检查输出时,可以沿同一条调用依次确认: -1. 原始 C/A/B regions 和 layouts; -2. dispatch 选择了哪个 variant; -3. A/B descriptors 的 major、swizzle、``ldo/sdo`` 和 slice offsets; -4. C datapath 与 TMEM offsets; -5. 生成的 MMA iteration 数量与 ``enable_input_d``; -6. cleanup 后还剩哪些直接 physical accesses。 +1. 原始 C/A/B regions 与 layouts; +2. dispatch 选中的 variant; +3. A/B descriptor 的 major、swizzle、``ldo/sdo`` 和 16-byte offsets; +4. C 的 ``TLane/TCol`` offsets 与 TMEM 地址; +5. MMA iteration 数量及 ``enable_input_d``; +6. cleanup 后生成的普通物理访问。 -这样可以把问题定位到 kernel schedule、layout contract、operator dispatcher -或后端 codegen,而不是只比较 TIRx 源码和最终 assembly。 +这样可以把问题定位到输入 schedule、layout 契约、算子 dispatcher 或后端 codegen 的具体阶段。 核心源码导航 ------------ - `dispatcher.py`_:variant registry、priority、predicate 与失败报告; -- `tile_primitive_dispatch.cc`_:scope/launch context、body replacement 与 - callbacks; -- `lower_tirx_cleanup.cc`_:``LayoutApplier``、``BufferOffsetRemover`` 与 - physical-offset materialization; -- `tcgen05 gemm dispatcher`_:从 operand regions/layouts 推导 descriptors、 - TMEM address、instruction tiling 和 ``T.ptx.tcgen05.mma``。 +- `tile_primitive_dispatch.cc`_:dispatch context、调用替换与 execution scope lowering; +- `lower_tirx_cleanup.cc`_:``LayoutApplier``、``BufferOffsetRemover`` 与物理地址物化; +- `tcgen05 gemm dispatcher`_:从操作数 regions/layouts 推导 descriptors、TMEM 地址、指令分块和 ``Tx.ptx.tcgen05.mma``。 .. _dispatcher.py: https://github.com/apache/tvm/blob/v0.26.0/python/tvm/tirx/operator/tile_primitive/dispatcher.py .. _tile_primitive_dispatch.cc: https://github.com/apache/tvm/blob/v0.26.0/src/tirx/transform/tile_primitive_dispatch.cc