资讯详情

CUTLASS CuTeDSL Blackwell GEMM 教程解析:从基础 FP16 Kernel 到 PDL 的程序化演进

📅 2026/9/15 20:58:08 | 华诺云谱 👁 阅读
CUTLASS CuTeDSL Blackwell GEMM 教程解析:从基础 FP16 Kernel 到 PDL 的程序化演进
CUTLASS CuTeDSL Blackwell GEMM 教程解析从基础 FP16 Kernel 到 PDL 的程序化演进【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass导读本文以 CUTLASS 仓库中的 Blackwell GEMM 教程系列examples/python/CuTeDSL/cute/blackwell/tutorial/tutorial_gemm为核心系统讲解如何用 CuTeDSLCUTLASS 的 Python DSL编写高性能 GEMM Kernel并逐步引入 Tensor Core、TMA、SMEM 布局、多级软件流水线、2CTA MMA、Warp 特化、持久化 Tile 调度、TMA 预取、Preferred Cluster 与程序化依赖启动PDL等优化手段。读完本文你将掌握该教程系列从fp16_gemm_0.py到fp16_gemm_6.py的完整演进脉络理解每项优化的动机、量化收益与适用场景并能在 8k×8k×8k 规模的 FP16 稠密 GEMM 与 NVFP4 块缩放 GEMM 上直接运行、修改和验证这些示例。教程定位与整体结构该目录是 CUTLASS 在 BlackwellSM100/SM100a架构上的一组教程示例展示如何利用 Tensor Core 编写高性能 GEMM通用矩阵乘法Kernel。按照 README.md 的概述教程覆盖了以下场景与技术点基础 FP16 GEMM 实现软件流水线Software Pipeline优化Tensor Core 利用线程/Warp/Block 三级并行。目录中实际包含 10 个示例脚本与 1 个公共工具模块文件主题fp16_gemm_0.py基础 FP16 GEMMTMA 加载、SMEM 布局、cutlass.range(..., prefetch_stages...)多级流水线fp16_gemm_1.py引入 2CTA MMA 与 TMA 多播multicastfp16_gemm_2.pyWarp 特化WSTMA / MMA / Epilogue 分派到不同 WarpEpilogue 使用 TMA Storefp16_gemm_3.py静态持久化 Tile 调度器StaticPersistentTileSchedulerfp16_gemm_3_1.py动态持久化 Tile 调度器ClcDynamicPersistentTileSchedulerfp16_gemm_4.pyPreferred / Fallback / Dynamic Cluster 支持fp16_gemm_5.pyTMA 预取TMA Prefetch优化fp16_gemm_6.py程序化依赖启动PDLnvfp4_gemm_0.pyNVFP4 块缩放block-scaled批处理 GEMM 基础实现nvfp4_gemm_1.pyNVFP4 块缩放 GEMM 2CTA 指令与 TMA 多播utils.py公共命令行解析、Tensor 构造、参考校验与基准测试工具从源码结构看示例之间存在清晰的依赖递进关系fp16_gemm_1在0的基础上做数据搬运与集群优化2引入 Warp 特化3/3_1引入持久化调度4引入集群形状的动态配置5引入 L2 预取6引入跨 Kernel 重叠NVFP4 系列则把同样的优化思路复用到块缩放 GEMM 上。示例 0CuTeDSL 基础 FP16 GEMM配置常量与运行方式fp16_gemm_0.py 定义了本示例的核心配置io_dtype cutlass.Float16 acc_dtype cutlass.Float32 mma_inst_shape_mnk (128, 256, 16) mma_tiler_mnk (128, 256, 64) threads_per_cta 128 # Pipeline stage configuration ab_stages 4 acc_stage 1io_dtype/acc_dtype输入为 FP16累加器为 FP32mma_inst_shape_mnk单条 MMA 指令形状 (128, 256, 16)mma_tiler_mnk单个 CTA 每次处理的 Tile 为 (128, 256, 64)threads_per_cta每 CTA 128 线程ab_stagesA/B 矩阵的流水线级数4 级acc_stage累加器缓冲级数1 级。运行命令对应 README.md 与脚本 docstringpython examples/python/CuTeDSL/cute/blackwell/tutorial/tutorial_gemm/fp16_gemm_0.py \ --mnk 8192,8192,8192脚本通过argparse接收--mnk逗号分隔的三个整数与--tolerance校验容差默认1e-01。约束条件是 m、n 必须能被 Tile 尺寸 (128, 256) 整除否则直接报错退出。运行前脚本会调用 CUDA Driver API 检查 GPU 是否存在fp16_gemm_0.py。使用cutlass.range(..., prefetch_stages...)构建多级软件流水线README 特别强调的cutlass.range(..., prefetch_stages...)是示例 0 的核心写法。在主循环中fp16_gemm_0.pynum_k_tiles cute.size(gA, mode[2]) if warp_idx 0: # Wait for a empty accumulator buffer acc_empty acc_producer.acquire_and_advance() for k_tile_idx in cutlass.range(num_k_tiles, prefetch_stagesab_stages - 2): # Issue TMA loads ab_empty ab_producer.acquire_and_advance() cute.copy( tma_atom_a, tAgA[(None, ab_empty.count)], tAsA[(None, ab_empty.index)], tma_bar_ptrab_empty.barrier, ) cute.copy( tma_atom_b, tBgB[(None, ab_empty.count)], tBsB[(None, ab_empty.index)], tma_bar_ptrab_empty.barrier, ) # Execute one K-block worth of MMA instructions ab_full ab_consumer.wait_and_advance() num_k_blocks cute.size(tCrA, mode[2]) for k_block_idx in cutlass.range_constexpr(num_k_blocks): k_block_coord (None, None, k_block_idx, ab_full.index) cute.gemm( tiled_mma, tCtAcc, tCrA[k_block_coord], tCrB[k_block_coord], tCtAcc, ) tiled_mma.set(tcgen05.Field.ACCUMULATE, True) # Signal that the A/B buffers have been consumed and are ready for the next load ab_full.release()cutlass.range的prefetch_stages参数会自动生成多级流水线的地址推进逻辑把原本需要手写的前导预取prologue、滚动加载与 MMA 重叠的样板代码boilerplate封装起来。这里prefetch_stagesab_stages - 2即 2流水线的生产者ab_producerTMA 加载与消费者ab_consumerMMA 计算通过pipeline.PipelineTmaUmma.create创建的 barrier 同步ab_stages4意味着同一时刻最多有 4 个 K-tile 的 A/B 数据在流水线中飞行从而隐藏 DRAM 延迟。TMA 与 SMEM 布局TMATensor Memory Access用于全局内存到共享内存的批量拷贝。示例 0 中共享内存张量sA/sB通过smem.allocate_tensor分配byte_alignment128并使用a_smem_layout.inner的 swizzle 模式fp16_gemm_0.pyTMA 原子通过cute.nvgpu.cpasync.CopyBulkTensorTileG2SOp(tcgen05.CtaGroup.ONE)构造再经make_tiled_tma_atom_A/make_tiled_tma_atom_B绑定全局张量与 SMEM 布局fp16_gemm_0.py内核启动前由 Warp 0 执行cpasync.prefetch_descriptor预取 TMA 描述符fp16_gemm_0.pyA/B 的 SMEM 布局由cutlass.utils.blackwell_helpers即sm100_utils的make_smem_layout_a/make_smem_layout_b生成二者在 host 端cute.jit函数host_function内计算并作为 kernel 参数传入。TMEM 分配与 Epilogue 子 Tile 化Blackwell 的 tcgen05 MMA 累加器位于张量内存TMEM中。示例 0 用utils.TmemAllocator分配 512 列 TMEM并通过NamedBarrier协调分配fp16_gemm_0.py。README 中提到的Tiling Epilogue to avoid bursty write out and reduce register pressure体现在 epilogue 部分累加结果被zipped_divide切成 4 个子 Tilesubtile_cnt 4每个子 Tile 通过tcgen05.Ld32x32bOp(tcgen05.Repetition.x64)从 TMEM 加载到寄存器tCrAcc转换回io_dtype后经cute.autovec_copy写回全局内存fp16_gemm_0.py。子 Tile 化有两个作用一是避免一次性爆发式写回造成的内存带宽尖峰二是降低单线程寄存器压力每线程仅需持有 64 个 fp32。正确性验证run_dense_gemm用 torch 生成 K-major 的随机张量值域 [-2, 2)通过cutlass.cute.runtime.from_dlpack包装为 CuTe 张量并标记动态布局调用 JIT 编译的host_function后用torch.einsum(mk,nk-mn, ...)计算参考结果再以torch.testing.assert_close校验fp16_gemm_0.py。这一构造输入 → 编译启动 → 参考校验的流程在后续所有示例中反复出现是 CuTeDSL 教程的标准验证范式。示例 12CTA MMA 与 TMA 多播fp16_gemm_1.py 在示例 0 基础上增加 2CTA MMA 与 TMA 多播其性能提升来源有两个方面1. 2CTA MMA 降低 B 张量 SMEM 占用从而增大 AB 流水线级数。脚本给出了详尽的 SMEM 容量分析单级 A 张量 SMEM 恒为128 × 64 × sizeof(fp16) 16KBB 张量在 1CTA 下为256 × 64 × sizeof(fp16) 32KB在 2CTA 下只需一半即 16KB。因此1CTA 最大 AB 级数227 // (16 32) 42CTA 最大 AB 级数227 // (16 16) 7。对应的延迟隐藏能力分别为512 × (4-1) 1.5K cycles与512 × (7-1) 3K cycles512 为一次 DRAM 往返的近似周期数227KB 为 SM100 每 CTA 的 SMEM 容量上限。这解释了为何fp16_gemm_1.py将ab_stages提升到 7。2. TMA 多播降低 L2 流量。对 (m, n) 的 cluster 形状单个 Tile 的 L2 流量为16KB / n 32KB / m2×1 cluster16KB/1 32KB/2 24KB4×4 cluster16KB/4 32KB/4 12KB。而若未启用 TMA 多播单个 Tile 的 L2 流量接近16KB 32KB 48KB。更大的流水线级数提供更强的延迟隐藏能力而多播缩短数据就绪时间二者需要针对延迟受限或内存吞吐受限的场景权衡。示例 1 的配置fp16_gemm_1.pycluster_shape_mnk (2, 1, 1) mma_inst_shape_mnk (256, 256, 16) mma_tiler_mnk (256, 256, 64) ab_stages 7 acc_stage 1注意 Tile 尺寸也升级为 (256, 256)m、n 须能被 256 整除。示例 2Warp 特化与 TMA Storefp16_gemm_2.py 在示例 1 基础上引入 Warp 特化Warp Specialization, WS把 TMA 加载、MMA 计算、Epilogue 三类任务分派到不同的 Warp 组。Warp 特化的核心思想是让同一 CTA 内的不同 Warp 各司其职并通过 barrier 通信。其收益来自 Warp 间的任务级并行DMA Warp 在完成当前 K-block 的 A/B 加载后立即开始下一个 K-block 的加载而 MMA Warp 同时在计算当前 K-block 的结果DRAM 延迟因此被隐藏非 WS 版本也能靠预取隐藏 DRAM 延迟但 WS 版本把不同类型指令TMA 加载、TMEM 分配、MMA放到不同 Warp 发射指令级并行更好。例如非 WS 版本中 TMEM 分配与 TMA 加载在同一个 Warp 串行执行TMA 加载必须等 TMEM 分配完成WS 版本中二者在不同 Warp 中可重叠。此外示例 2 的 Epilogue 改用 TMA Store 写回全局内存这需要两步先把结果从寄存器写入 SMEM再由 SMEM 通过 TMA Store 写回全局。Epilogue 继续采用子 Tile 化既降低 epilogue 的 SMEM 占用又能把下一个子 Tile 的st.shared与当前子 Tile 的 TMA Store 重叠隐藏st.shared延迟。文档头还给出了 WS 版本的适用性分析fp16_gemm_2.py大 MMA Tile 且 AB 级数足够时主循环性能两者相近WS 的收益主要来自 prologue/epilogue——因此 K 维度较小时 WS 优势明显小 MMA Tile 时 WS 主循环也可能更快因为 MMA 每条指令附带的 ALU 准备工作占比更高WS 把 TMA 与 MMA 的 ALU 分散到不同 WarpMMA Warp 中 ALU 指令更少MMA 发射更高效。配置上示例 2 新增threads_in_epilogue 128epilogue 线程数、epi_stages 2epilogue 流水级数并通过use_2cta_instrs开关控制是否使用 2CTA 指令fp16_gemm_2.py。示例 3 / 3_1持久化 Tile 调度器静态持久化调度示例 3fp16_gemm_3.py 使用StaticPersistentTileScheduler以固定、确定的顺序把 Tile 分配给 CTA适合划分良好的工作负载。持久化 cluster 在整个 kernel 执行期间驻留在 GPU 上并处理多个 Tile从而隐藏 prologue 与 epilogue 开销问题规模越大、Tile 数越多prologue/epilogue 在不同 Tile 之间被隐藏的效果越好主循环相对较短、prologue/epilogue 占比较高时收益尤其明显。文档也明确提示其局限静态调度在部分 SM 资源不可用时容易出现负载不均衡这正是下一个示例引入动态调度器的原因。示例 3 将acc_stages提升到 2并新增scheduler_type utils.StaticPersistentTileSchedulerfp16_gemm_3.py。动态持久化调度示例 3_1fp16_gemm_3_1.py 改用ClcDynamicPersistentTileScheduler比静态调度器更灵活能更好地处理工作负载不均衡。动态调度器在SharedStorage中额外分配了 CLCCluster Launch Control相关的 barrier 与响应缓冲区# Only for CLC Dynamic Scheduler clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 2] clc_response_align_bytes num_clc_response_bytes clc_response: cute.struct.Align[ cute.struct.MemRange[cutlass.Int32, 4], clc_response_align_bytes, ]其中num_clc_response_bytes 16响应大小为 4B × 4 个元素。脚本用use_clc_dynamic_scheduler开关在动态与静态调度器之间切换fp16_gemm_3_1.py。与示例 3 不同这里 kernel 参数中张量改由TmaInfotma_a/tma_b/tma_c封装传递mA_mkl等直接从TmaInfo.tma_tensor取出fp16_gemm_3_1.py。示例 4Dynamic Cluster 与 Preferred Clusterfp16_gemm_4.py 演示 CuTeDSL 对 Preferred Cluster 与 Dynamic Cluster 的支持。Compute Capability 9.0 引入 Thread Block Clusters 层级更大的 cluster 形状能获得更高的 TMA 多播效率但因量化效应可能导致 SM 占用不佳——例如 18 SM 的 GPU 上使用 2×2 cluster 只占用 16 个 SM留下 2 个 SM 空闲。Compute Capability 10.0 起可同时指定两个 clusterpreferred cluster 与 fallback cluster示例 3_1 中即可额外启动一个 2×1 cluster 来利用空闲的 2 个 SM。术语定义与约束条件文档头原文要点Static cluster编译期指定 cluster 形状Dynamic cluster运行期由 host 设置 cluster 形状Preferred clusterkernel 可同时携带 preferred 与 fallback 两种 cluster 形状启动约束preferred cluster 的深度Z 维度必须与 fallback 相同fallback 形状必须能整除 preferred 形状preferred 形状必须能整除启动网格形状。示例 4 的关键配置fp16_gemm_4.pyfallback_cluster_shape_mnk (2, 1, 1) if use_2cta_instrs else (1, 1, 1) preferred_cluster_shape_mnk (2, 4, 1) if use_2cta_instrs else (1, 1, 1)即 preferred 为 2×4、fallback 为 2×1二者 Z 维度相同、fallback 能整除 preferred。用户可通过 kernel 启动参数指定 preferred 与 fallback 形状。示例 5TMA Prefetchfp16_gemm_5.py 演示 TMA 预取优化用cute.prefetch()提前把数据从 DRAM 带入 L2 缓存对内存受限memory-bound工作负载尤其有益。预取分为两个阶段初始阶段Initial Phase在 TMA 加载循环开始前先把前prefetch_dist个 K-tile 预取到 L2为后续 TMA 拷贝预热缓存滚动阶段Rolling PhaseTMA 加载循环的每次迭代中在发出当前 K-tile 的 TMA 拷贝后预取前方prefetch_dist处的 K-tile保持 L2 持续预热。与示例 3_1 的三个关键差异文档头原文新增cute.prefetch()调用主 TMA 循环前有初始预取循环加载循环内有滚动预取。示例 5 还调整了配置以适配预取场景fp16_gemm_5.pycluster 为 2×2、mma_inst_shape_mnk (256, 64, 16)、mma_tiler_mnk (256, 64, 64)、ab_stages 10——更细的 N 维度与更高流水级数配合预取进一步吸收 DRAM 延迟。示例 6程序化依赖启动PDLfp16_gemm_6.py 演示 Programmatic Dependent LaunchPDL一种允许同一 stream 内前后两个 kernel 重叠执行的机制。典型场景是反量化 GEMM第二个操作的 GEMM 的 B 操作数是第一个操作的输出。当 GEMM 主循环只占总时间的一小部分、prologue/epilogue 主导时可以提前启动第二个 GEMM 的 prologue 并与反量化 kernel 重叠。最紧的依赖约束是GEMM 必须在反量化 kernel 完成把反量化后的 B 写入全局内存之后才开始主循环。启用 PDL 需要做两件事文档头原文在 kernel 中插入griddepcontrol.launch_dependents与griddepcontrol.wait指令启动 kernel 时设置 PDL launch 属性。griddepcontrol.launch_dependents与griddepcontrol.wait提供对 PDL 中 kernel 执行的细粒度控制一旦所有线程块执行了launch_dependents依赖的 kernel 就可以机会性地提前启动。文档头给出实测示例对--mnk 256,8192,128PDL 相对无 PDL 的加速比最高可达 1.16×。该示例还引入了反量化 kernel 的配置quant_dtype cutlass.Int8、dequant_elements_per_thread 128、threads_in_dequant 64并建议使用较小的问题规模运行默认--mnk 256,8192,128。NVFP4 块缩放 GEMM 系列基础版本nvfp4_gemm_0.pynvfp4_gemm_0.py 是 NVFP4 块缩放批处理 GEMM 的入门实现输入 A/B 为Float4E2M1FNNVFP4缩放因子 SFA/SFB 为Float8E4M3FN输出 C 为Float16缩放因子向量大小为 16每 16 个元素共享一个缩放因子并携带批维度 l。核心配置mma_tiler_mn (128, 256) mma_inst_shape_k 64 ab_dtype cutlass.Float4E2M1FN sf_dtype cutlass.Float8E4M3FN c_dtype cutlass.Float16 sf_vec_size 16脚本以类Sm100BlockScaledDenseGemmKernel封装 kernel内部通过utils.get_smem_capacity_in_bytes(sm_100)查询 SMEM 容量num_ab_stage 4、num_acc_stage 1、num_tmem_alloc_cols 512nvfp4_gemm_0.py。运行命令python examples/python/CuTeDSL/cute/blackwell/tutorial/tutorial_gemm/nvfp4_gemm_0.py \ --mnkl 8192,8192,8192,1 --do_benchmark约束条件文档头原文m、n、k 须能被 Tile (128, 256, 256) 整除缩放因子向量大小为 16A/B 在 k 维连续C 在 n 维连续A/B 为Float4E2M1FNSFA/SFB 为Float8E4M3FN。工具模块 utils.py 提供了create_parser--mnkl、--tolerance、--do_benchmark参数、to_blocked把参考缩放因子转为 cuBLAS 块缩放布局以及通用的run函数它构造随机 NVFP4 输入与缩放因子张量通过cvt_sf_MKL_to_M32x4xrm_K4xrk_L一个cute.jit辅助函数负责把 MKL 布局的缩放因子转换为 MMA 所需的 32×4×rm 布局生成 CuTe 格式缩放因子用cute.compile编译并启动 kernel再用torch._scaled_mm计算参考结果并assert_close校验开启--do_benchmark时调用cute.testing.benchmark统计执行时间、FLOPS 与带宽utils.py。2CTA 与 TMA 多播版本nvfp4_gemm_1.pynvfp4_gemm_1.py 把示例 1 的优化思路复用到 NVFP4 块缩放 GEMM同样给出量化的 SMEM 分析单级 A128 × 256 × sizeof(float4) 16KB单级 sfA128 × (256/16) × sizeof(float8) 2KB单级 sfB256 × (256/16) × sizeof(float8) 4KBB 在 1CTA 下为256 × 256 × sizeof(float4) 32KB2CTA 下只需一半16KB。因此最大 AB 级数为1CTA227 // (163224) 42CTA227 // (161624) 5延迟隐藏能力分别为512 × (4-1) 1.5K cycles与512 × (5-1) 2K cycles。TMA 多播的 L2 流量分析与示例 1 一致2×1 cluster 为 24KB/tile4×4 cluster 为 12KB/tile未多播约 48KB/tile。该示例配置mma_tiler_mn (256, 256)、cluster_shape_mnk (2, 1, 1)nvfp4_gemm_1.pym、n、k 须能被 (256, 256, 256) 整除。贯穿系列的通用工程模式从源码可以归纳出 CuTeDSL Blackwell GEMM 教程反复使用的工程范式便于读者举一反三cute.struct定义SharedStorage在 Python 侧声明 SMEM 布局mbar barrier、TMEM 持有缓冲、CLC 响应区等配合SmemAllocator.allocate使用参见各示例的SharedStorage类cute.kernel定义 device kernelcute.jit定义 host 端 JIT 函数host 函数负责构造 Tiled MMAtcgen05.MmaF16BF16Op、SMEM 布局与 TMA 原子计算网格形状cute.ceil_div((*c.layout.shape, 1), mma_tiler_mnk[:2])后调用kernel(...).launch(grid..., block...)fp16_gemm_0.py两段式流水线同步PipelineTmaUmma负责 A/B 的 TMA→SMEM→MMA 流水PipelineUmmaAsync负责 TMEM 累加器空/满状态的同步二者通过acquire_and_advance/wait_and_advance/release/commit协作fp16_gemm_0.pyTMEM 生命周期管理Warp 0 分配 →tmem.wait_for_alloc()全 CTA 同步 →retrieve_ptr获取累加器指针 → epilogue 完成后relinquish_alloc_permitpipeline.synctmem.freefp16_gemm_0.py统一的验证与基准from_dlpack包装 torch 张量 →cute.compile/no_cacheTrue编译 → 与 torch 参考实现assert_closecute.testing.benchmark输出执行时间、PFLOPS 与带宽。总结与选型建议回顾整个教程系列可以得到一条清晰的性能优化决策路径从fp16_gemm_0.py基础 TMA 多级流水线 Tile 化 Epilogue出发先确认正确性与基线性能若高 SM 频率下 DRAM 延迟成为瓶颈采用fp16_gemm_1.py的 2CTA MMA节省 B 的 SMEM、增大 AB 级数与 TMA 多播降低 L2 流量若 K 较小或 MMA Tile 较小、prologue/epilogue 占比高采用fp16_gemm_2.py的 Warp 特化与 TMA Store若 Tile 数众多用fp16_gemm_3.py/3_1.py的持久化 Tile 调度器隐藏 prologue/epilogue动态调度器对负载不均衡更稳健SM 量化导致占用不佳时用fp16_gemm_4.py的 Preferred/Fallback Cluster 补齐空闲 SM内存受限时用fp16_gemm_5.py的 TMA 预取把 DRAM 延迟藏进 L2多 kernel 流水线场景如反量化GEMM用fp16_gemm_6.py的 PDL 实现 kernel 间重叠。NVFP4 系列则把 2CTA 与 TMA 多播的思路迁移到块缩放 GEMM验证了这些优化对Float4E2M1FN输入同样有效。所有示例均附带可复现的运行命令、约束条件与参考校验读者可以直接修改--mnk/--mnkl、ab_stages、cluster 形状等参数观察不同配置下的性能与正确性变化快速建立对 Blackwell GEMM 调优的直觉。【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。