diff --git a/zh/tirx_guide/arch/index.rst b/zh/tirx_guide/arch/index.rst index db9145de..f2ff1fc6 100644 --- a/zh/tirx_guide/arch/index.rst +++ b/zh/tirx_guide/arch/index.rst @@ -20,9 +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 d5ecd121..4cc1405d 100644 --- a/zh/tirx_guide/arch/lowering_pipeline.rst +++ b/zh/tirx_guide/arch/lowering_pipeline.rst @@ -15,205 +15,142 @@ 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 做一类特定的转换、检查或标注。 - -TIRx 的完整 pass 顺序定义在 Apache TVM 源码中的 `python/tvm/tirx/compilation_pipeline.py -`_。 +用 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,最终交给目标后端生成代码。 -整体编译路径 +完整编译路径 ------------ -``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.tile.* TilePrimitiveCall + ├─ Buffer + TileLayout + └─ device / CTA / warpgroup / warp / thread scopes + │ + │ BindTarget + ▼ + LowerTIRx + ├─ TilePrimitiveDispatch:Tx.tile.* → target-specific TIR / Tx.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 同时服务于两个阶段:``TilePrimitiveDispatch`` 较早读取它来查找算子实现,最终 code generator 再读取它来生成目标代码。``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 顺序 -------------------------------- + 高层 tile primitive 路径 -下表按照实际执行顺序列出 ``tirx_pipeline`` 中的 19 个步骤。ABI 是函数之间的调用约定;表中的 ABI passes 负责把普通 TIR 函数改造成 runtime 能够调用的形式。``PassContext`` 是控制编译选项的配置对象:公共子表达式消除可以关闭,向量化和循环展开的行为也可以通过它调整。 + Tx.tile.gemm_async + │ TilePrimitiveDispatch + ▼ + Tx.ptx.tcgen05.mma + │ + ▼ + tirx_pipeline 后续 passes -.. list-table:: - :header-rows: 1 - :widths: 6 24 24 46 + 直接 PTX 路径 - * - # - - 类别 - - 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 保存 + Tx.ptx.tcgen05.mma + │ + ▼ + tirx_pipeline 后续 passes -Host 与 Device 的后续处理 -------------------------- +因此,直接 PTX 从对应 tile primitive 的算子级选择、检查和参数推导之后接入,并继续经过 ``tirx_pipeline`` 中余下的 passes。两条路径的详细比较见 :ref:`chap_tirx_tile_layout_lowering`。 -上面列出的 19 个步骤组成 ``tirx_pipeline``。这条模块级 pipeline -结束后,``tvm.compile`` 会根据函数类型分别执行 finalization: +``LowerTIRx`` 的边界 +------------------------------ -- **host**:``LowerTVMBuiltin`` 处理 ``tvm_*`` builtins,``LowerIntrin`` - 处理面向具体 target 的 intrinsics。 -- **device**:``LowerWarpMemory`` 将 warp-scoped buffers 转换为 - shuffles,随后执行 ``StmtSimplify`` 和 ``LowerIntrin``。 +默认情况下,``LowerTIRx`` 包含两个 transformation passes: -``LowerTIRx`` 的内部组成 ------------------------- +.. code-block:: text -``LowerTIRx`` 主要完成两个任务:为 tile-level 操作选择具体实现,以及把逻辑数据布局转换成实际的内存索引。它的核心转换由下面两个 passes 组成,定义在 Apache TVM 源码中的 -`src/tirx/transform/lower_tirx.cc -`_: + LowerTIRx = Sequential([ + TilePrimitiveDispatch, + LowerTIRxCleanup, + ]) + +``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 - LowerTIRx = Sequential([ TilePrimitiveDispatch, LowerTIRxCleanup ]) + bx = Tx.cta_id([grid_x]) → bx = blockIdx.x + tx = Tx.thread_id([block_x]) → tx = threadIdx.x -- **``TilePrimitiveDispatch``** 为 tile 操作选择具体实现。TIRx 中的 ``copy``、 - ``gemm``、``reduction`` 等操作以 ``TilePrimitiveCall`` 表示;这个 pass 根据 - backend 选择对应实现。它还会把 ``T.cta_id``、``T.thread_id`` 等抽象的执行范围编号转换成 kernel launch 参数和线程绑定。 -- **``LowerTIRxCleanup``** 将逻辑坐标转换为物理索引。它把支持的逻辑 layout 应用到 buffer access 上,使后续 passes 可以直接处理具体的索引表达式。 +``LowerTIRxCleanup`` 随后对仍然存在的直接 ``BufferLoad`` / ``BufferStore`` 应用 memory layout,展平相应 buffers,清除已消费的 layout metadata,并移除 ``tirx.buffer_offset`` wrappers。必须先 dispatch 后 cleanup:算子实现需要在 metadata 消失前读取完整的 region、shape、dtype、scope 和 layout。 -完成 ``LowerTIRx`` 后,tile 操作已经换成选定的底层实现,逻辑 layout 也已经落实为物理索引,``T.cta_id`` 和 ``T.thread_id`` 等抽象编号则变成了线程绑定。此时仍可能保留 thread-binding loops 和 TIRx 特有的 loop annotations;后续的 -``LowerTIRxOpaque`` 会规范化这些结构,再由 ``tirx.transform.FlattenBuffer`` -展平普通 TIR 中的 buffer access。 +``LowerTIRx`` 成功结束后: -一个简单 Kernel 的编译过程 --------------------------- +- ``TilePrimitiveCall`` 已被具体实现替换; +- 抽象 scope IDs 已变成 launch parameters、``Bind`` 和 thread bindings; +- layout 已被算子 lowering 消费,或被 cleanup 物化为后端能处理的地址; +- ``Tx.ptx.*`` 等 target intrinsics 仍可存在; +- thread-binding loops 和 TIRx-specific loop annotations 仍可能存在,随后由 + ``LowerTIRxOpaque`` 规范化; +- host/device split、ABI lowering 和最终 code generation 留给后续阶段。 -下面用一个 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 源码使用抽象的线程编号。** +下面用一个执行逐元素计算的 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]) - B[bx * 256 + tx] = A[bx * 256 + tx] * T.float32(2.0) + from tvm.script import tirx as Tx -``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 提供的抽象编号。 + @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] * Tx.float32(2.0) -**2. ``LowerTIRx`` 将抽象编号转换为 TIR 线程绑定。** 它把 ``bx`` 和 ``tx`` -分别绑定到 ``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) - B[bx * 256 + tx] = A[bx * 256 + tx] * T.float32(2.0) - -这里仍然是 TIR,还不是 CUDA 源码。这段代码只保留了关键映射,不是编译器输出的完整 IR;下一节会给出打印完整结果的命令。 + with Tx.launch_thread("blockIdx.x", 4) as bx: + tx = Tx.launch_thread("threadIdx.x", 256) + i = bx * 256 + tx + B[i] = A[i] * Tx.float32(2.0) -**3. 后续 passes 拆分 host/device,并生成 CUDA。** 编译开始时只有一个 TIRx -函数。``LowerTIRx`` 生成线程绑定和 device region 后,``SplitHostDevice`` 将其拆成两个 TIR 函数(PrimFunc): +这里展示的是省略 buffer declarations 后的 TIR 结构摘要。完整 printer output 还会包含声明等细节,CUDA 源码则由后续 codegen 生成。``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 +159,211 @@ 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 由输入 IR 提供 + * - 类型合法化 + - BF16/FP8 compute legalization、BF16/FP8 storage legalization、 + final ``LowerIntrin`` + - 将当前 target/backend 需要降级表示的 compute、storage 和 intrinsic + 改写为受支持形式;具备原生表示能力的 target 保留原表示 + * - 校验、模块和 ABI lowering + - ``VerifyMemory``、``AnnotateEntryFunc``、``SplitHostDevice``、 + ``LowerIket``、``MakePackedAPI`` + - 验证设备计算处于 thread environment,抽取 device kernel,降低 + kernel launch,并生成 runtime 可调用的 packed ABI + +``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。 + +.. code-block:: text + + 一个包含 device region 的 PrimFunc + │ + │ SplitHostDevice + ▼ + host launcher device PrimFunc + ├─ 准备调用参数 ├─ thread_extent + ├─ grid/block launch parameters ├─ physical buffer accesses + └─ 调用 device kernel └─ Tx.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++;部分 ``Tx.ptx.*`` intrinsic 会被打印为 inline PTX 或 helper code。之后 NVRTC/NVCC/ptxas 等 CUDA toolchain 组件继续产生可加载的 PTX 或 binary。``LowerIntrin`` 的产物仍由随后运行的 CUDA code generator 和 toolchain 继续处理。 + +性能决策的职责边界 +------------------ + +默认 TIRx pipeline 的责任边界如下: + +.. 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 分解为硬件指令 + * - ``LowerTIRx`` 后续 passes + - 局部化简、index narrowing、显式或配置驱动的 vectorize/unroll,以及 + dtype、module 和 ABI 合法化 + * - CUDA backend 与 toolchain + - 生成 target code,进行更底层的 peephole、register allocation 和组装 -这里不需要边界判断,因为 ``4 * 256`` 恰好等于 1,024。处理一般长度 ``N`` -时,需要向上取整得到 CTA 数量,并在 kernel 中判断 ``i < N``。 +默认 pipeline 的自动优化集中在局部程序变换。Tile sizes、跨算子 fusion、software pipeline、warp specialization 和主要 layout 等整体 schedule 决策,通常由 kernel 作者或上层生成器提供。Dispatcher 则根据给定 shape 将一个 tile operation 映射成一组硬件指令;全程序 schedule search、cost-model selection 和 autotuning 属于更上层的生成与搜索系统。 -检查中间 IR 与生成代码 ----------------------- +检查完整流水线 +-------------- -为了查看中间 IR,可以只运行完整 pipeline 最前面的几步,然后停下来打印结果。下面先把 ``scale`` 以全局名 ``main`` 放入 ``IRModule``。CUDA target 指定 GPU -端生成 CUDA,``with_host("llvm")`` 则指定 CPU 端生成 LLVM 代码。 -``BindTarget`` 将这组 target 信息写入 module,随后只运行 ``LowerTIRx``: +前面的 ``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 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`` 时执行 18 步序列,后续 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、折叠其余 + unit loops,并规范化 loop pragmas + * - 5 + - ``FlattenBuffer`` + - 展平剩余 ``BufferLoad`` / ``BufferStore`` + * - 6 + - ``BF16ComputeLegalize`` + - 在需要 fallback 的 target 上将 BF16 compute 提升至 ``float32`` 并改写 + * - 7 + - ``NarrowDataType(32)`` + - 能够证明安全时,将 index/loop expressions 缩窄到 32 bits + * - 8 + - ``VectorizeLoop`` + - Lower 已标记的 vectorized loops;标量模式下按标量 loops 处理 + * - 9 + - ``UnrollLoop`` + - 展开显式 ``Tx.unroll``,并执行配置或 pragma 允许的自动展开 + * - 10 + - ``StmtSimplify`` + - 在 vectorize/unroll 暴露常量后再次化简 + * - 11 + - ``CommonSubexprElim`` + - 执行公共子表达式消除,是否运行由 ``tir.disable_cse_tir`` 控制 + * - 12 + - ``FP8ComputeLegalize`` + - 在需要 fallback 的 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`` 表示 + +判断默认的 aggressive auto-vectorization 或 auto-unrolling 行为时,需要同时检查 pass 名称、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..36aaaf10 --- /dev/null +++ b/zh/tirx_guide/arch/tile_primitive_layout_lowering.rst @@ -0,0 +1,468 @@ +.. 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` 给出了完整的编译流水线。本页进入 ``LowerTIRx`` 内部,沿一条 Blackwell ``Tx.tile.gemm_async`` 说明 tile primitive 和 layout 怎样逐步变成 target intrinsic 与物理地址。 + +``Tx.tile.gemm_async(C, A, B)`` 记录一次逻辑 tile GEMM。它给出操作数、区域和配置,具体使用哪些硬件指令、每个操作数怎样编码,则由 ``LowerTIRx`` 结合 target 与 layout 确定。这里的 layout 是一张从 **逻辑坐标** 到 **物理坐标** 的映射:生产者用它摆放数据,消费者用它找到同一份数据。 + +``LowerTIRx`` 由两个顺序执行的 pass 组成: + +.. code-block:: text + + TilePrimitiveCall + Buffer layouts + target + │ + ▼ + TilePrimitiveDispatch + 选择实现,生成 Tx.ptx.* 与地址表达式 + │ + ▼ + LowerTIRxCleanup + 物化剩余普通访问,清除 layout metadata + │ + ▼ + 包含 Tx.ptx.* 与物理 Buffer 访问的 PrimFunc + +前一个 pass 读取算子级语义,后一个 pass 处理仍留在普通 ``BufferLoad``、``BufferStore`` 和指针访问中的 layout。同一个 buffer 的 layout 可以依次服务于这两个阶段。 + +先读懂 Layout +------------- + +一个逻辑 tile 可以落在普通内存、多个线程的局部存储或 TMEM 中。三类位置需要三种物理坐标: + +.. list-table:: + :header-rows: 1 + :widths: 22 35 43 + + * - 存储位置 + - Layout 映射 + - 物理坐标的含义 + * - GMEM / SMEM + - ``(i, j) → m`` + - ``m`` 是一个线性存储位置,也可以包含 padding 或 swizzle + * - 每线程局部存储 + - ``(i, j) → (thread axis, m)`` + - thread axis 指定持有元素的线程,``m`` 指定该线程内的局部位置 + * - Blackwell TMEM + - ``(i, j) → (TLane, TCol)`` + - 两个坐标共同指定 TMEM 中的硬件位置 + +这些映射都由 ``TileLayout`` 表达。完整的 layout 代数、``S[...]`` 语法与组合规则见 :ref:`chap_tirx_layout_api`;这里关注它们在 lowering 中提供的信息。 + +普通内存:逻辑下标映射到地址 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +省略 ``layout=`` 时,buffer API 默认创建连续行优先(row-major)layout。例如 shape 为 ``(4, 8)`` 的 buffer 使用: + +.. code-block:: text + + TileLayout(S[(4, 8)]) + + B[i, j] → B_flat[i * 8 + j] + +这个公式与普通二维连续数组一致。layout metadata 会在编译前半程保留,因此 dispatcher 仍可对它执行 slice、匹配与检查。显式 ``layout=None`` 会让该字段保持为空,普通访问随后沿 shape/stride 规则展开。 + +物理排列也可以带 padding 或 swizzle。例如: + +.. code-block:: text + + TileLayout(S[(4, 8) : (16, 1)]) + + B[i, j] → B_flat[i * 16 + j] + +这里每一行占 16 个元素,逻辑 shape 仍是 ``4×8``。对应的 backing allocation 需要覆盖 ``layout.span() = 56`` 个元素。``ComposeLayout`` 还可以把 XOR、shift 和 mask 组成 shared-memory swizzle。 + +每线程局部存储:逻辑下标映射到持有者和局部位置 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +下面的 distributed layout 把第 ``m`` 行交给 workgroup 中的第 ``m`` 个线程,并把 ``n`` 作为该线程的局部位置: + +.. code-block:: python + + Dreg_wg = Dreg.view( + 128, + N, + layout=TileLayout(S[(128, N) : (1@tid_in_wg, 1)]), + ) + +.. code-block:: text + + Dreg_wg[m, n] → { tid_in_wg: m, local slot: n } + +``tid_in_wg`` 表达元素归属,默认 memory axis 记录该线程内的局部位置。CUDA toolchain 随后为这些局部值分配具体物理寄存器。 + +TMEM:逻辑下标映射到两个硬件坐标 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Blackwell TMEM 使用 ``TLane`` 和 ``TCol``: + +.. code-block:: text + + TileLayout(S[(128, N) : (1@TLane, 1@TCol)]) + + C[m, n] → { TLane: m, TCol: n } + +``TLane`` 表示 TMEM lane row,``TCol`` 表示 TMEM column。二者共同组成 tcgen05 指令使用的 TMEM 地址。这里的 ``TLane`` 属于存储坐标;CUDA ``lane_id`` 表示当前执行线程,两者含义不同。 + +Layout 从哪里来 +~~~~~~~~~~~~~~~~ + +TIRx kernel 中的 layout 通常来自三处: + +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`` 等接口写出目标映射。 + +这里的自动化包含默认 layout 的构造,以及 lowering 根据既有 layout 推导指令参数。tile size、swizzle mode、线程分工和性能 layout 的选择通常由 kernel 作者、上层生成器或搜索系统完成。``LowerTIRx`` 接收这些已经确定的映射,再完成局部推导与合法性检查。 + +沿一条 ``Tx.tile.gemm_async`` 看 lowering +------------------------------------------- + +下面抽取 :ref:`chap_tirx_primer` 中单-tile GEMM 的核心语句。SMEM/TMEM allocation、barrier 和等待代码在这里省略: + +.. 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) + + C = Tx.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 Tx.ptx.elect_sync(): + Tx.tile.gemm_async( + C[:, :BLK_N], + Asmem[:, :], + Bsmem[:, :], + accum=False, + dispatch="tcgen05", + cta_group=1, + ) + +这一条调用向 dispatcher 提供了四组信息: + +.. list-table:: + :header-rows: 1 + :widths: 30 70 + + * - 输入 + - 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 + +选择实现 +~~~~~~~~ + +Python 前端把 ``Tx.tile.gemm_async`` 保存为 ``TilePrimitiveCall``。候选实现以 ``(operator, target kind)`` 为键注册,并带有 variant 名称、优先级和适用条件。 + +显式 ``dispatch="tcgen05"`` 会筛选出对应 variant;省略 ``dispatch=`` 时,dispatcher 按优先级检查各候选的适用条件。选中的实现返回一个 ``PrimFunc``,其函数体替换原来的 ``TilePrimitiveCall``。整个过程发生在编译期,依据 target、操作数 region、layout、dtype 和 config 做规则匹配。 + +其他 copy、reduce 和同步类 tile primitive 也经过同一套选择流程。每个实现自行定义要读取哪些 layout 坐标、接受哪些 shape,以及最终生成哪些 target intrinsics。 + +切出本次调用的 Layout +~~~~~~~~~~~~~~~~~~~~~~ + +调用参数是 ``C[:, :BLK_N]``、``Asmem[:, :]`` 和 ``Bsmem[:, :]`` 这样的 ``BufferRegion``。Dispatcher 先对各 buffer layout 执行 ``slice``,把 region 起点折入 layout offset,再通过 ``canonicalize`` 得到便于匹配的等价形式。 + +tcgen05 实现随后检查: + +- 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 与候选实现的硬件约定一致后,lowering 才继续生成 descriptor 和指令。 + +A/B Layout 推导 matrix descriptor +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +tcgen05 从 SMEM 读取 A/B 时,硬件通过 matrix descriptor 解释 shared-memory 地址。Dispatcher 将 sliced layout 与受支持的 K-major 或 MN-major swizzle atom 匹配,并得到: + +.. code-block:: text + + major mode + swizzle mode + leading-dimension offset (ldo) + stride-dimension offset (sdo) + 当前指令 tile 相对 region 原点的 16-byte offset + +``ldo``、``sdo`` 和 swizzle 等字段可在 dispatch 时确定;SMEM base address 属于运行时值。生成的 TIR 会调用 ``Tx.ptx.tcgen05.encode_matrix_descriptor``,把运行时地址与这些字段编码到一个局部 ``uint64`` descriptor 中。后续各个 MMA iteration 再加入以 16 bytes 为单位的 tile offset。 + +这也解释了 producer 与 consumer 为何要共享 layout。前面的 copy 按 A/B layout 把元素写入 swizzled SMEM,GEMM dispatcher 从同一份 layout 生成 descriptor,tcgen05 因而按照相同的字节排列读回元素。 + +C Layout 推导 TMEM 地址 +~~~~~~~~~~~~~~~~~~~~~~~ + +示例中 C region 的映射为: + +.. code-block:: text + + C[m, n] → { TLane: m, TCol: n } + +对 layout 做 slice 后,dispatcher 分别取得 ``TLane`` offset 和 ``TCol`` offset。它们与 ``allocated_addr`` 一起传给 ``Tx.cuda.get_tmem_addr``,形成 tcgen05 使用的 TMEM 目标地址: + +.. code-block:: text + + get_tmem_addr(allocated_addr, TLane offset, TCol offset) + +当 C region 从非零行或非零列开始时,对应偏移会进入这两个坐标。Dispatcher 同时验证 sliced layout 与所选 accumulator datapath 一致。 + +Shape 与 dtype 决定指令分解 +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +示例表示一次 fp16 ``128×128×64`` GEMM。tcgen05 的这个 dense 路径每个 K step 处理 16 个 fp16 元素,因此 K 方向被分成 4 次 MMA。省略具体 TIR 语法后,dispatch 结果的结构如下: + +.. code-block:: text + + # 结构示意;尖括号中的内容代表生成的 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), + ... + ) + +这里的 ``desc_a`` 和 ``desc_b`` 需要运行时 SMEM 地址,所以 descriptor encoding 保留在生成的 TIR 中。dense ``desc_i`` 只依赖指令 shape、dtype、major mode 与 ``cta_group``,dispatcher 会把它折叠为编译期 ``uint32`` 常量。 + +``Tx.unroll(4)`` 会由后续 ``UnrollLoop`` 展开,``StmtSimplify`` 再化简每次迭代中的常量表达式。指令数量与操作数分块来自 tile primitive 实现,循环展开和局部化简由流水线后段完成。 + +Readback 使用另一组 Layout +~~~~~~~~~~~~~~~~~~~~~~~~~~ + +GEMM 将结果写入 TMEM 后,readback 是另一条独立的 tile primitive: + +.. code-block:: python + + 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]) + +它同时读取 C 的 ``TLane/TCol`` layout 和 ``Dreg_wg`` 的 distributed layout,选择相应的 ``tcgen05.ld`` form,并把 TMEM 元素送到指定线程的局部位置。整个数据流为: + +.. code-block:: text + + 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) + +Layout 在这里充当 producer 与 consumer 之间的物理映射契约;tile primitive 负责把这份契约翻译成具体硬件操作。 + +``LowerTIRxCleanup`` 处理剩余 Layout +------------------------------------ + +``TilePrimitiveDispatch`` 已将全部 ``TilePrimitiveCall`` 展开。普通 ``BufferLoad``、``BufferStore`` 和指针访问仍可能引用带 layout 的 buffer,``LowerTIRxCleanup`` 随后处理这些访问: + +.. code-block:: text + + 创建或附着 Layout + │ + ▼ + Buffer view / region slice + │ + ▼ + TilePrimitiveDispatch 读取并展开算子 + │ + ▼ + LowerTIRxCleanup 映射剩余普通访问 + │ + ▼ + 重建物理 Buffer,清除 layout metadata + │ + ▼ + 后续 TIR passes 与 target codegen + +对普通 memory layout,``LayoutApplier`` 的核心过程可以写成: + +.. code-block:: text + + logical indices + │ + ▼ + layout.canonicalize().apply(indices, shape) + │ + ▼ + 一个 symbolic memory coordinate + │ + ▼ + physical buffer access + +这里的 symbolic coordinate 可以包含运行时循环变量、``threadIdx`` 或函数参数。后续 ``FlattenBuffer`` 会把 buffer 的 ``elem_offset`` 合入线性 index,形成最终地址语义。``BufferOffsetRemover`` 则消除 ``tirx.buffer_offset(BufferLoad)`` wrapper,使其中的 offset 与物理访问保持一致。 + +普通指针访问需要 layout 最终产生一个 memory coordinate,通常命名为 ``m``。带有 ``laneid``、``tid_in_wg``、``TLane`` 或 ``TCol`` 的映射表达额外的硬件坐标;它们应先由理解这些坐标的 tile primitive 消费,或通过 ``.view()`` 与 ``.local()`` 转成当前线程的存储视图。若这类坐标以普通 ``BufferLoad`` 留到 cleanup,编译器会给出诊断。 + +``LowerTIRx`` 完成后,所有 ``TilePrimitiveCall`` 都已展开,layout metadata 已被 dispatcher 读取或由 cleanup 物化,``Tx.ptx.*`` 等 target intrinsics 则继续进入后续 TIR passes 与 CUDA codegen。 + +高层 Tile Primitive 与直接 PTX +------------------------------- + +直接写 ``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: 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 推导 + - 作者调用低级编码与地址接口 + * - 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 +-------------------------------- + +需要定位问题时,可以分别打印 authored TIRx、dispatch 结果和完整 ``LowerTIRx`` 结果: + +.. 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()) + +示例使用 Blackwell ``tcgen05``,因此 target 设为 ``sm_100a``。检查输出时,可以沿同一条调用依次确认: + +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 后生成的普通物理访问。 + +这样可以把问题定位到输入 schedule、layout 契约、算子 dispatcher 或后端 codegen 的具体阶段。 + +核心源码导航 +------------ + +- `dispatcher.py`_:variant registry、priority、predicate 与失败报告; +- `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 +.. _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