- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
本篇技术指南以 TVM TIRx 的 tile primitivegemm为核心,系统讲解同步矩阵乘法如何被下放(lower)为全展开的 warp 协作mma.sync.aligned.m16n8k{16,8}指令嵌套,覆盖接受条件、演示程序、碎片(fragment)排布算法、生成的 TIRx IR 与 CUDA 代码。读者读完将掌握 TIRx 中寄存器级 GEMM 的调用约定、m16n8k16/m16n8k8指令选择机制、PTX 操作数枚举顺序以及如何通过测试用例验证数值正确性,并理解它与 Blackwelltcgen05.mma异步路径(gemm_async)的分工边界。
一、gemm在 TIRx 中的定位
TVM TIRx 将内核中的硬件级操作建模为一组可派发的tile primitive(见 tile_primitives.rst)。一次 primitive 调用在 IR 中记录为一个未解析的TilePrimitiveCall节点,编译期由TilePrimitiveDispatch阶段根据 primitive 名称、执行范围(thread / warp / warpgroup / cta)、操作数布局、目标后端和可选的显式 hint 选择具体下放实现,并替换为原生 IR。矩阵乘法族包含两个成员:
gemm:同步路径,全部操作数驻留寄存器,下放为 warp 级mma.sync;gemm_async:异步路径,走 Blackwelltcgen05.mma,A/B 通常驻留共享内存、累加器驻留张量内存(见 gemm_async.rst)。
本文聚焦同步gemm的 CUDA 下放实现,其源码位于 mm_m16n8k_.py,对应的完整测试位于 test_gemm_mma_m16n8k_.py。
语义定义
gemm计算:
D = alpha·A@B + beta·C它被下放为一个全展开(fully-unrolled)的 warp 协作mma.sync.aligned.m16n8k{16,8}指令嵌套。A/B 碎片与 C/D 累加器全部位于寄存器(local作用域)——调用方需要先把 A/B 通过(通常)ldmatrix从共享内存装载为寄存器碎片(参见 copy/ldstmatrix.rst)。派发器将 M/N/K 切分为m16n8k原子块,每个输出 tile 发射一条mma,并在 K 维上就地累加(in-place accumulate)。
二、What it accepts:派发门槛与接受条件
该变体在 mm_m16n8k_.py 中以register_dispatch("gemm", "cuda", variant="mma.m16n8k*", priority=10, ...)注册,携带两个谓词full_active_lanes与no_replica。原文档给出的注册骨架如下:
# register_dispatch("gemm", "cuda", priority=10, when=[ predicate("full_active_lanes", _full_active_lanes), # complete warp(s), un-narrowed predicate("no_replica", _no_replica), # no broadcast axes on D/A/B/C # ]) # in the impl: for buf, name in ((D, "D"), (A, "A"), (B, "B"), (C, "C")): if buf.scope() != "local": fail(f"gemm mma requires {name} in register (local) scope, got {buf.scope()}")| Property | Requirement |
|---|---|
| target / scope / priority | cuda;priority10。谓词不设作用域白名单:要求sctx.intra中出现的每个轴都是完整、零偏移的laneid/wid_in_wg/warpid轴。这接纳常规的 warp / warpgroup / CTA 调用点,拒绝簇轴等未识别轴。mma.sync仍是 warp 协作的,因此调用方使用 warp 或由完整 warp 组成的更宽作用域 |
| operand scope | A、B、C、D 全部位于寄存器(local);共享内存操作数会让派发fail(先用 ldmatrix 装载) |
| no replica | D/A/B/C 均不得携带 broadcast/replica 轴(_no_replica) |
| shape | M % 16 == 0、N % 8 == 0、K % 8 == 0。派发器先尝试m16n8k16再尝试m16n8k8;被选中的指令必须能精确切分所有操作数布局 |
| dtype | A 与 B 同为float16或同为bfloat16;C 与 D 为float32 |
| alpha / beta | alpha == 1.0;beta ∈ {0.0, 1.0}(0 →D = A@B;1 →D = A@B + C) |
源码级的接受条件详解
从实现源码可以更精确地还原每一步检查(对应 mm_m16n8k_.py):
- 作用域检查:遍历
(D, "D"), (A, "A"), (B, "B"), (C, "C"),任一缓冲的scope()不是local即调用fail(...)拒绝该调用。这正是"纯寄存器路径"约束的落地处——调用方负责预先装载碎片(典型路径是copy → ldstmatrix)。 _full_active_lanes(mm_m16n8k_.py):mma.sync.aligned对每个活动线程是集体操作,任何窄化intra轴的if包裹都会使.aligned指令行为未定义,因此要求每个intra轴偏移为 0 且 extent 完整:laneid=32、wid_in_wg=4(warpgroup 场景)、warpid=warps-per-CTA(CTA 场景,从启动配置的threadIdx.xextent 除以 32 推导)。出现任何其他轴(例如簇作用域的cta_id)即被拒绝。_no_replica(mm_m16n8k_.py):检查 D/A/B/C 四个操作数布局的replica字段,任一存在 broadcast 轴即拒绝。- alpha/beta 约束:
mma.sync原生计算D = A·B + C(无标量缩放),因此只支持alpha=1,beta只接受 0 或 1——beta决定 C 是否作为累加器初值(1 →c_ptr=C,0 →c_ptr=0)。源码通过Analyzer().simplify()求值常量标量,非 1.0 的alpha与非 {0,1} 的beta都会触发fail(mm_m16n8k_.py)。 - 维度一致性:
transpose_A/transpose_B只描述输入的逻辑朝向,实现归一化为标准形A=[M,K]、B=[K,N](D/C 恒为[M,N]),并断言A.K == B.K、D == (M,N)、C == (M,N)。 - 指令表驱动:
MMA_INSTRUCTIONS是一个可扩展表(mm_m16n8k_.py),当前含四个条目:m16n8k16.bf16、m16n8k16.f16、m16n8k8.bf16、m16n8k8.f16,全部为k_pack=2(沿 K 每寄存器打包 2 个 16 位元素)。新增指令只需追加条目,可行性检查保持通用。
派发失败时的行为也值得注意:任何谓词不通过或实现内部调用fail(reason),都会抛出一个DispatchFail,派发器记录拒绝原因并继续尝试下一个候选,若全部失败则在最终RuntimeError中汇总所有拒绝原因(机制见 dispatcher.py 与 tile_dispatch.rst)。
三、演示程序:单 warp 完成 16×16 @ 16×8
原文档给出一个最小可运行的演示:单个 warp 计算D[16,8] = A[16,16] @ B[16,8](fp16 输入、f32 累加),恰好对应一个m16n8k16原子(取自 test_gemm_mma_m16n8k_.py 的数值测试):
from tvm.tirx.layout import S, TileLayout, laneid D_FRAG = TileLayout(S[(2, 8, 4, 2) : (2, 4 @ laneid, 1 @ laneid, 1)]) A_FRAG_K8 = TileLayout(S[(2, 8, 4, 2) : (2, 4 @ laneid, 1 @ laneid, 1)]) B_FRAG_K8 = TileLayout(S[(4, 2, 8) : (1 @ laneid, 1, 4 @ laneid)]) A_FRAG = A_FRAG_K8.tile_to([16, 16], [16, 8]); B_FRAG = B_FRAG_K8.tile_to([16, 8], [8, 8]) @Tx.prim_func def gemm(A_ptr: Tx.handle, B_ptr: Tx.handle, D_ptr: Tx.handle): A_g = Tx.match_buffer(A_ptr, (16, 16), "float16"); B_g = Tx.match_buffer(B_ptr, (16, 8), "float16") D_g = Tx.match_buffer(D_ptr, (16, 8), "float32") Tx.device_entry(); Tx.cta_id([1]); Tx.warp_id([1]); lane = Tx.lane_id([32]) A_f = Tx.alloc_buffer((16, 16), "float16", scope="local", layout=A_FRAG) B_f = Tx.alloc_buffer((16, 8), "float16", scope="local", layout=B_FRAG) D_f = Tx.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) A_reg = A_f.local(8) # stage A into the lane's 8 regs for s in Tx.unroll(8): kp, rM, kHi = s % 2, (s // 2) % 2, s // 4 A_reg[s] = A_g[lane // 4 + 8 * rM, 2 * (lane % 4) + kp + 8 * kHi] B_reg = B_f.local(4) # stage B into the lane's 4 regs for s in Tx.unroll(4): kp, kHi = s % 2, s // 2 B_reg[s] = B_g[2 * (lane % 4) + kp + 8 * kHi, lane // 4] Tx.tile.warp.gemm(D_f, A_f, B_f, D_f, transpose_A=False, transpose_B=False, alpha=1.0, beta=0.0) D_reg = D_f.local(4) # write the 4 result regs out for s in Tx.unroll(4): rN, rM = s % 2, s // 2 D_g[lane // 4 + 8 * rM, 2 * (lane % 4) + rN] = D_reg[s]逐段解读
- 碎片布局定义:
S[... : ...]表示共享(shard)迭代器对逻辑维度的映射。D_FRAG的(2, 8, 4, 2) : (2, 4 @ laneid, 1 @ laneid, 1)表示每个线程拥有 4 个 f32 累加寄存器,坐标为(rM, rN),映射到逻辑坐标M = lane//4 + 8·rM、N = 2·(lane%4) + rN——这正是 PTX ISA m16n8 累加器寄存器映射c_id = 2·rM + rN。A_FRAG/B_FRAG则由A_FRAG_K8/B_FRAG_K8通过tile_to沿 K 堆叠得到(k16 = 两个 k8 沿 K 拼接)。 - 装载碎片:
Tx.unroll循环把全局内存数据按mma的物理寄存器顺序解码进每个 lane 的寄存器槽:A 的物理槽序为s = 4·kHi + 2·rM + kp,B 为s = 2·kHi + kp。这里不能用整块T.copy,因为 per-thread 轴无法按坐标匹配。 - 发起 GEMM:
Tx.tile.warp.gemm(D_f, A_f, B_f, D_f, ...)——注意 C 位置传入的是D_f本身(beta=0 时无所谓,beta=1 时意味着"把 D 当作累加器初值",即纯就地累加形态)。 - 写回:按
D的物理寄存器顺序c_id = 2·rM + rN解码回逻辑(M, N)坐标写回全局内存。
四、Algorithm:三步下放算法
1. Tile 与 fragment-group
派发器把每个操作数的布局切片(slice)到其 region,然后对每个候选指令(先m16n8k16后m16n8k8)尝试把操作数子布局(D_M, D_N, A_M, A_K, B_K, B_N, C_*)组合进固定的 m16n8k 框架:以 D 的 M 锚定 A/C,以 D 的 N 锚定 B/C,以 A 的 K 锚定 B 的 K。第一个形状与 warp 切分都匹配的指令胜出。
源码中对应的关键步骤(mm_m16n8k_.py):
_slice_group先把每个缓冲的布局slice到其 region 再group成 2D 缓冲序形状(A 为(M,K)/(K,M),B 为(K,N)/(N,K),C/D 为(M,N)),并按 transpose 标志交换轴序;布局不可切分时干净地拒绝。- 每个逻辑维做锚定对齐(anchor-align):
_align(DM, AM, ...)、_align(DM, CM, ...)、_align(DN, BN, ...)、_align(DN, CN, ...)、_align(AK, BK, ...),使每个共享逻辑维在所有操作数中以相同方式分解(follower 按锚的迭代 extent 分组并按锚的规范序重排,保留自身 stride)。 - 每维的线程/内存 region 长度由锚一次定死(
_region_totals),follower 复用该切分——注释明确指出,例如 B 的 N 其"寄存器槽"实际是 lane 迭代器,若按自身迭代器推导会错误地报告内存长度 1。
2. 推导寄存器布局
每个操作数得到一个经仲裁的每-lane 视图,同时保留mma.sync期望的PTX 操作数枚举顺序(原始物理寄存器序):
| 操作数 | 逻辑视图(shape 维序) | 物理寄存器序(默认local视图) |
|---|---|---|
| D/C | [Mo, No, rM, rN](4 个 f32) | [Mo, No, rM, rN] |
| A | [Mo, Ko, rM, kHi, k_pack] | [Mo, Ko, kHi, rM, k_pack] |
| B | [Ko, No, kHi, k_pack] | [Ko, No, kHi, k_pack] |
(Mo/No/Ko是 warp 级 tile 索引,rM/rN是累加器寄存器索引,kHi是高 K 寄存器组索引,k_pack是最内层沿 K 的连续打包。源码见 mm_m16n8k_.py。)
实现中_frag_group按指令的固定 lane 切分把每个子布局分组为碎片形状,并对每个切出的组验证"必须是单个迭代器、线程/内存轴类型正确、stride 匹配"。m16n8 的 lane 切分是:M 方向 8 个 lane(g = lane>>2,stride 4),N/K 方向 4 个 lane(t = lane&3,stride 1)。K 的内存尾部为[kHi, k_pack],其中k_pack是最内层(stride-1)连续打包、kHi = inst.k // (4·k_pack)是高 K 寄存器组数——m16n8k16的kHi=2,m16n8k8的kHi=1(这是 k8 路径必须特殊处理 extent-1 组的原因)。
随后跨操作数校验 warp 切分一致性:M.to(D/A/C 的组 0)必须逐元素匹配,N.to(D/B/C 的组 0)匹配,K.to(A/B 的组 0)匹配,确保同一逻辑块在三个操作数中落在同一 warp 上;任一不匹配即跳过该指令。
3. 发射展开的嵌套
初始化 D(beta==1时从 C 拷贝,否则清零),然后在 K 上就地累加,每个(m, n)tile 一条mma:
for m in Tx.unroll(M_tiles): for n in Tx.unroll(N_tiles): for rM, rN in ...: d_local[m, n, rM, rN] = c_local[...] if use_c else Tx.float32(0) for k in Tx.unroll(K_tiles): d_regs = [d_local[m, n, rM, rN] for rM in range(2) for rN in range(2)] # 4 f32 a_regs = [a_words[m, k, rM, kHi, 0] for kHi in range(n_kHi) for rM in range(2)] b_regs = [b_words[k, n, kHi, 0] for kHi in range(n_kHi)] mma_chain = (f"mma.sync.aligned.{shape_str}.row.col" f".f32.{a_elem}.{b_elem}.f32") Tx.ptxmma_chain # d = a·b + d要点:
- 循环边界全部是编译期常量,
T.unroll由UnrollLooppass 完全展开,因此本地缓冲索引解析为静态寄存器槽——mma的寄存器操作数必须是常量。 - A/B 碎片以每 b32 两个 16 位元素打包,因此通过
uint32视图(a_local.view("uint32")/b_local.view("uint32"))到达指令;PTX 链中的元素格式(f16/bf16/f32)与 TVM dtype 通过_PTX_ELEM映射转换(mm_m16n8k_.py)。 - 寄存器计数按指令推导而非硬编码:D/C 累加器
rM = inst.m//8、rN = inst.n//4(k16 与 k8 都是 4 个 f32);A 的 b32 数为rM + 2·kHi(k16 为 4,k8 为 2);B 的 b32 数为kHi(k16 为 2,k8 为 1)。 mma是d = a·b + c形态:D 累加器每个输出 tile 初始化一次(beta==1拷 C、beta==0清零),之后每个 K 步以c = d就地累加,得到统一的一条 mma 形态。
五、生成的 TIRx IR 与 CUDA
单个 16×8×16 tile 下放为一条mma(4 个 D 寄存器、4 个 A 寄存器、2 个 B 寄存器),生成的 TIRx IR 为:
Tx.ptx"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32"生成的 CUDA 内联汇编为:
"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 " "{%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {%0, %1, %2, %3};"其中累加器{%0..%3}同时是 C 输入和 D 输出(就地累加);{%4..%7}是 A 的四个b32寄存器,{%8, %9}是 B 的两个。该路径已在sm_100a上验证(D == A@B,fp16 容差内)。从 test_gemm_mma_m16n8k_.py 还可以看到:完整流水线(UnrollLoop + CUDA codegen)中每条mma以__device__helper 形式发射一次,helper 调用次数恰好等于Mt·Nt·Kt(ptx_mma_sync_aligned_m16n8k{kinst}_row_col的出现次数减 1)。
六、How inputs change the algorithm:输入如何改变下放
| input | effect |
|---|---|
| input dtype | float16→…f32.f16.f16.f32;bfloat16→…f32.bf16.bf16.f32(寄存器计数不变——每个b32仍是 2 个元素) |
| K instruction | k16→ A 4 个b32/ B 2 个b32;k8→ A 2 个 / B 1 个(mma.…m16n8k8.…) |
| M / N / K extents | 设置M_tiles/N_tiles/K_tiles展开循环计数(每个(m, n)一条mma,K 就地累加) |
| beta | 0→ D 零初始化;1→ D 从 C 初始化(mma本身完全一致) |
| transpose_A / transpose_B | 在 shape 与布局匹配前转置逻辑 A 或 B region;变换后的操作数仍须适配所选m16n8k框架 |
| operand scope | A/B必须是寄存器碎片;共享内存操作数使派发fail(先用ldstmatrix装载) |
测试覆盖(从源码结构可以确认)
test_gemm_mma_m16n8k_.py 对上述每类输入都提供了对应测试:
- 注册验证:
test_cuda_gemm_mma_variant_is_registered断言("tirx.tile.gemm", "cuda")的调度表包含"mma.m16n8k*"变体。 - 降级与 beta 语义:
test_cuda_gemm_mma_lowers_to_mma_sync断言 beta=0 时出现T.float32(0清零且 D 寄存器槽为d_local[0..3]、A 为a_words[0..3]、B 为b_words[0..1];test_cuda_gemm_mma_accumulates_c_when_beta_one断言 beta=1 时出现c_local[且不再清零。 - 拒绝路径:
test_cuda_gemm_mma_rejects_nonunit_alpha(alpha=2.0)、test_cuda_gemm_mma_rejects_fractional_beta(beta=0.5)、test_cuda_gemm_mma_rejects_unsupported_dtype(f16 累加、混合 A/B 输入、tf32、int8 四类签名全部拒绝)——这些测试确保实现"拒绝而不是发射错误的 mma"。 - 分块与数值:
_TILED_SHAPES×_TILED_MODES的笛卡尔积(9 种分块 × 4 种 (dtype, beta) 组合)覆盖单 tile、单维多 tile、全维 tile、M=64、k8 的kHi==1以及 K=24(16 不整除 24)等情形;test_cuda_gemm_mma_numerical_tiled在真机上用assert_allclose校验D = A@B (+ C)。 - 转置:
test_cuda_gemm_mma_numerical_transpose覆盖(transpose_A, transpose_B)的全部四种组合 × 两种 dtype,转置碎片布局通过group+permute_by_groups从 K-major 基布局推导得到,物理寄存器顺序不变。
这些测试中大部分断言只需 CPU 侧的LowerTIRxtransform 即可运行,仅数值校验需要requires_cuda真机(sm_100a实测验证)。
七、与gemm_async的分工边界
原文档将同步gemm定位为gemm_async的对照面(交叉引用见 gemm_async.rst)。两者核心差异可概括为:
- 操作数驻留:
gemm的 A/B/C/D 全在寄存器;gemm_async的 B(通常还有 A)在共享内存并由 64 位矩阵描述符命名,A 可走 tensor-memory 操作数路径,累加器在张量内存。 - 执行与同步:
gemm是 warp 集体的同步mma.sync;gemm_async由单线程(或经elect_sync选举)发起tcgen05.mma异步执行,调用方用tcgen05.commit+ mbarrier 等待完成。 - 精度与指令族:
gemm当前支持 f16/bf16 输入 + f32 累加(m16n8k16/k8);gemm_async支持 f16/bf16/fp8/fp4(含块缩放 SFA/SFB)、cta_group1/2、weight_stationary等多种模式。
对于不需要异步流水、且希望完全用寄存器承载数据的小 tile 计算,同步gemm是直接而高效的路径。
八、延伸阅读
- tile_primitives.rst——tile primitive 调用约定(
Tx.tile.warp.<name>作用域绑定)、TilePrimitiveCall字段与各变体消费的config键。 - arch/tile_dispatch.rst——
TilePrimitiveDispatchpass 的流水线、变体选择排序规则(priority 降序、变体名升序)与拒绝原因汇总。 - operator/tile_primitive/dispatcher.py——
register_dispatch/predicate/fail的实现契约(实现必须返回PrimFunc或抛DispatchFail)。 - tile_primitives/copy/ldstmatrix.rst——装载 A/B 寄存器碎片的标准前置步骤(
ldmatrix/stmatrix的 m8n8 碎片几何)。 - tile_primitives/gemm_async.rst——Blackwell
tcgen05.mma异步路径对照。 - layout.rst 与 api/layout.rst——
TileLayout模型与S[...]/tile_to等碎片布局操作。
- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
相关推荐
TVM TIRx 寄存器拷贝路径(vec_auto reg)深度解析:从 Tile 布局到 `ld.shared.v4.u32` 的完整降级链路
TVM TIRx 寄存器拷贝路径(vec_auto reg)深度解析:从 Tile 布局到 ld.shared.v4.u32 的完整降级链路 本指南深入剖析 T
模型编译深度学习推理引擎TVM TIRx tcgen05 张量内存与寄存器异步拷贝:copy_async tmem<->local 变体(tcgen05.ld/st)深度解析
TVM TIRx tcgen05 张量内存与寄存器异步拷贝:copy_async tmem< local 变体(tcgen05.ld/st)深度解析 本篇技术指
模型编译深度学习推理引擎TVM TIRx 降级流水线深度解析:从 Tile 原语到 CUDA Kernel 的完整编译路径
TVM TIRx 降级流水线深度解析:从 Tile 原语到 CUDA Kernel 的完整编译路径 导读 :本文围绕 TIRx 降级流水线文档 https://
模型编译深度学习推理引擎
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考