TVM TIRx 同步 GEMM 深度解析:warp 级 mma.sync 的寄存器碎片排布与完整下放链路
2026/9/23 5:22:15 网站建设 项目流程
  • 模型编译
  • 深度学习
  • 推理引擎

【免费下载链接】tvm

Open Machine Learning Compiler Framework

项目地址:https://gitcode.com/gh_mirrors/tv/tvm
点击查看免费下载

本篇技术指南以 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_lanesno_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()}")
PropertyRequirement
target / scope / prioritycuda;priority10。谓词不设作用域白名单:要求sctx.intra中出现的每个轴都是完整、零偏移的laneid/wid_in_wg/warpid轴。这接纳常规的 warp / warpgroup / CTA 调用点,拒绝簇轴等未识别轴。mma.sync仍是 warp 协作的,因此调用方使用 warp 或由完整 warp 组成的更宽作用域
operand scopeA、B、C、D 全部位于寄存器local);共享内存操作数会让派发fail(先用 ldmatrix 装载)
no replicaD/A/B/C 均不得携带 broadcast/replica 轴(_no_replica
shapeM % 16 == 0N % 8 == 0K % 8 == 0。派发器先尝试m16n8k16再尝试m16n8k8;被选中的指令必须能精确切分所有操作数布局
dtypeA 与 B 同为float16或同为bfloat16;C 与 D 为float32
alpha / betaalpha == 1.0beta ∈ {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=32wid_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=1beta只接受 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.KD == (M,N)C == (M,N)
  • 指令表驱动MMA_INSTRUCTIONS是一个可扩展表(mm_m16n8k_.py),当前含四个条目:m16n8k16.bf16m16n8k16.f16m16n8k8.bf16m16n8k8.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]

逐段解读

  1. 碎片布局定义S[... : ...]表示共享(shard)迭代器对逻辑维度的映射。D_FRAG(2, 8, 4, 2) : (2, 4 @ laneid, 1 @ laneid, 1)表示每个线程拥有 4 个 f32 累加寄存器,坐标为(rM, rN),映射到逻辑坐标M = lane//4 + 8·rMN = 2·(lane%4) + rN——这正是 PTX ISA m16n8 累加器寄存器映射c_id = 2·rM + rNA_FRAG/B_FRAG则由A_FRAG_K8/B_FRAG_K8通过tile_to沿 K 堆叠得到(k16 = 两个 k8 沿 K 拼接)。
  2. 装载碎片Tx.unroll循环把全局内存数据按mma的物理寄存器顺序解码进每个 lane 的寄存器槽:A 的物理槽序为s = 4·kHi + 2·rM + kp,B 为s = 2·kHi + kp。这里不能用整块T.copy,因为 per-thread 轴无法按坐标匹配。
  3. 发起 GEMMTx.tile.warp.gemm(D_f, A_f, B_f, D_f, ...)——注意 C 位置传入的是D_f本身(beta=0 时无所谓,beta=1 时意味着"把 D 当作累加器初值",即纯就地累加形态)。
  4. 写回:按D的物理寄存器顺序c_id = 2·rM + rN解码回逻辑(M, N)坐标写回全局内存。

四、Algorithm:三步下放算法

1. Tile 与 fragment-group

派发器把每个操作数的布局切片(slice)到其 region,然后对每个候选指令(先m16n8k16m16n8k8)尝试把操作数子布局(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 寄存器组数——m16n8k16kHi=2m16n8k8kHi=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.unrollUnrollLooppass 完全展开,因此本地缓冲索引解析为静态寄存器槽——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//8rN = inst.n//4(k16 与 k8 都是 4 个 f32);A 的 b32 数为rM + 2·kHi(k16 为 4,k8 为 2);B 的 b32 数为kHi(k16 为 2,k8 为 1)。
  • mmad = 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·Ktptx_mma_sync_aligned_m16n8k{kinst}_row_col的出现次数减 1)。

六、How inputs change the algorithm:输入如何改变下放

inputeffect
input dtypefloat16…f32.f16.f16.f32bfloat16…f32.bf16.bf16.f32(寄存器计数不变——每个b32仍是 2 个元素)
K instructionk16→ A 4 个b32/ B 2 个b32k8→ A 2 个 / B 1 个(mma.…m16n8k8.…
M / N / K extents设置M_tiles/N_tiles/K_tiles展开循环计数(每个(m, n)一条mma,K 就地累加)
beta0→ D 零初始化;1→ D 从 C 初始化(mma本身完全一致)
transpose_A / transpose_B在 shape 与布局匹配前转置逻辑 A 或 B region;变换后的操作数仍须适配所选m16n8k框架
operand scopeA/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.syncgemm_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——Blackwelltcgen05.mma异步路径对照。
  • layout.rst 与 api/layout.rst——TileLayout模型与S[...]/tile_to等碎片布局操作。
  • 模型编译
  • 深度学习
  • 推理引擎

【免费下载链接】tvm

Open Machine Learning Compiler Framework

项目地址:https://gitcode.com/gh_mirrors/tv/tvm
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询