
模型编译深度学习推理引擎【免费下载链接】tvmOpen 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.syncgemm_async异步路径走 Blackwelltcgen05.mmaA/B 通常驻留共享内存、累加器驻留张量内存见 gemm_async.rst。本文聚焦同步gemm的 CUDA 下放实现其源码位于 mm_m16n8k_.py对应的完整测试位于 test_gemm_mma_m16n8k_.py。语义定义gemm计算D alpha·AB 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, variantmma.m16n8k*, priority10, ...)注册携带两个谓词full_active_lanes与no_replica。原文档给出的注册骨架如下# register_dispatch(gemm, cuda, priority10, 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(fgemm mma requires {name} in register (local) scope, got {buf.scope()})PropertyRequirementtarget / scope / prioritycudapriority10。谓词不设作用域白名单要求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_replicashapeM % 16 0、N % 8 0、K % 8 0。派发器先尝试m16n8k16再尝试m16n8k8被选中的指令必须能精确切分所有操作数布局dtypeA 与 B 同为float16或同为bfloat16C 与 D 为float32alpha / betaalpha 1.0beta ∈ {0.0, 1.0}0 →D AB1 →D AB C源码级的接受条件详解从实现源码可以更精确地还原每一步检查对应 mm_m16n8k_.py作用域检查遍历(D, D), (A, A), (B, B), (C, C)任一缓冲的scope()不是local即调用fail(...)拒绝该调用。这正是纯寄存器路径约束的落地处——调用方负责预先装载碎片典型路径是copy → ldstmatrix。_full_active_lanesmm_m16n8k_.pymma.sync.aligned对每个活动线程是集体操作任何窄化intra轴的if包裹都会使.aligned指令行为未定义因此要求每个intra轴偏移为 0 且 extent 完整laneid32、wid_in_wg4warpgroup 场景、warpidwarps-per-CTACTA 场景从启动配置的threadIdx.xextent 除以 32 推导。出现任何其他轴例如簇作用域的cta_id即被拒绝。_no_replicamm_m16n8k_.py检查 D/A/B/C 四个操作数布局的replica字段任一存在 broadcast 轴即拒绝。alpha/beta 约束mma.sync原生计算D A·B C无标量缩放因此只支持alpha1beta只接受 0 或 1——beta决定 C 是否作为累加器初值1 →c_ptrC0 →c_ptr0。源码通过Analyzer().simplify()求值常量标量非 1.0 的alpha与非 {0,1} 的beta都会触发failmm_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_pack2沿 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, scopelocal, layoutA_FRAG) B_f Tx.alloc_buffer((16, 8), float16, scopelocal, layoutB_FRAG) D_f Tx.alloc_buffer((16, 8), float32, scopelocal, layoutD_FRAG) A_reg A_f.local(8) # stage A into the lanes 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 lanes 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_AFalse, transpose_BFalse, alpha1.0, beta0.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 kpB 为s 2·kHi kp。这里不能用整块T.copy因为 per-thread 轴无法按坐标匹配。发起 GEMMTx.tile.warp.gemm(D_f, A_f, B_f, D_f, ...)——注意 C 位置传入的是D_f本身beta0 时无所谓beta1 时意味着把 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_totalsfollower 复用该切分——注释明确指出例如 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 个 laneg lane2stride 4N/K 方向 4 个 lanet lane3stride 1。K 的内存尾部为[kHi, k_pack]其中k_pack是最内层stride-1连续打包、kHi inst.k // (4·k_pack)是高 K 寄存器组数——m16n8k16的kHi2m16n8k8的kHi1这是 k8 路径必须特殊处理 extent-1 组的原因。随后跨操作数校验 warp 切分一致性M.toD/A/C 的组 0必须逐元素匹配N.toD/B/C 的组 0匹配K.toA/B 的组 0匹配确保同一逻辑块在三个操作数中落在同一 warp 上任一不匹配即跳过该指令。3. 发射展开的嵌套初始化 Dbeta1时从 C 拷贝否则清零然后在 K 上就地累加每个(m, n)tile 一条mmafor 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 (fmma.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//4k16 与 k8 都是 4 个 f32A 的 b32 数为rM 2·kHik16 为 4k8 为 2B 的 b32 数为kHik16 为 2k8 为 1。mma是d a·b c形态D 累加器每个输出 tile 初始化一次beta1拷 C、beta0清零之后每个 K 步以c d就地累加得到统一的一条 mma 形态。五、生成的 TIRx IR 与 CUDA单个 16×8×16 tile 下放为一条mma4 个 D 寄存器、4 个 A 寄存器、2 个 B 寄存器生成的 TIRx IR 为Tx.ptxmma.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 ABfp16 容差内。从 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输入如何改变下放inputeffectinput 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)一条mmaK 就地累加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断言 beta0 时出现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断言 beta1 时出现c_local[且不再清零。拒绝路径test_cuda_gemm_mma_rejects_nonunit_alphaalpha2.0、test_cuda_gemm_mma_rejects_fractional_betabeta0.5、test_cuda_gemm_mma_rejects_unsupported_dtypef16 累加、混合 A/B 输入、tf32、int8 四类签名全部拒绝——这些测试确保实现拒绝而不是发射错误的 mma。分块与数值_TILED_SHAPES×_TILED_MODES的笛卡尔积9 种分块 × 4 种 (dtype, beta) 组合覆盖单 tile、单维多 tile、全维 tile、M64、k8 的kHi1以及 K2416 不整除 24等情形test_cuda_gemm_mma_numerical_tiled在真机上用assert_allclose校验D AB ( C)。转置test_cuda_gemm_mma_numerical_transpose覆盖(transpose_A, transpose_B)的全部四种组合 × 两种 dtype转置碎片布局通过grouppermute_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/k8gemm_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等碎片布局操作。赞分享模型编译深度学习推理引擎【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址https://gitcode.com/gh_mirrors/tv/tvm点击查看免费下载相关推荐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),仅供参考