拓冰建站拓冰建站
首页 / 资讯中心 / 正文

CUTLASS 高效 GEMM 实现指南:分层分块结构、软件流水线与 Warp 特化设计

CUTLASS 高效 GEMM 实现指南分层分块结构、软件流水线与 Warp 特化设计【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass导读本指南以 CUTLASS 官方文档 efficient_gemm.md 为主体系统讲解 CUTLASS 如何将 GEMM矩阵乘加映射到 NVIDIA GPU 的线程块threadblock、线程束warp与线程thread三级并发模型覆盖分层分块结构、Epilogue 阶段、软件流水线、Threadblock Rasterization、SplitK/SlicedK 并行归约以及 Hopper 起引入的 Warp 特化Warp Specialization三种持久化内核设计。读完本文你将理解 CUTLASS GEMM 内核的核心调度结构、各层 tiling 参数如ThreadblockShape、WarpShape、MmaShape的作用并能依据问题规模选择合适的内核设计策略。分层分块结构Hierarchical Structure计算矩阵乘法的基本三重循环for i / for j / for k可以被分块blocking与平铺tiling以匹配硬件中的并发度、内存局部性与并行编程模型。CUTLASS 中 GEMM 映射到 NVIDIA GPU 的结构由下面这段循环嵌套示意for (int cta_n 0; cta_n GemmN; cta_n CtaTileN) { // for each threadblock_y } threadblock-level concurrency for (int cta_m 0; cta_m GemmM; cta_m CtaTileM) { // for each threadblock_x } for (int cta_k 0; cta_k GemmK; cta_k CtaTileK) { // GEMM mainloop - no unrolling // - one iteration of this loop is one stage // for (int warp_n 0; warp_n CtaTileN; warp_n WarpTileN) { // for each warp_y } warp-level parallelism for (int warp_m 0; warp_m CtaTileM; warp_m WarpTileM) { // for each warp_x } // for (int warp_k 0; warp_k CtaTileK; warp_k WarpTileK) { // fully unroll across CtaTileK // - one iteration of this loop is one k Group // for (int mma_k 0; mma_k WarpTileK; mma_k MmaK) { // for each mma instruction } instruction-level parallelism for (int mma_n 0; mma_n WarpTileN; mma_n MmaN) { // for each mma instruction } for (int mma_m 0; mma_m WarpTileM; mma_m MmaM) { // for each mma instruction } // mma_instruction(d, a, b, c); // TensorCore matrix computation } // for mma_m } // for mma_n } // for mma_k } // for warp_k } // for warp_m } // for warp_n } // for cta_k } // for cta_m } // for cta_n这个平铺后的循环嵌套针对三个层面的并发进行优化线程块threadblock之间的并发线程束warp之间的并发CUDA 核心与 Tensor Core即mma指令级的并发。同时它利用了存储层次中的局部性**共享内存shared memory**中的局部性**寄存器registers**中的局部性。下图展示了数据在该结构内的流动方式即 CUTLASS 所体现的分层 GEMM 计算从左到右每一级 stage 都对应一层嵌套的 tiling分别对应 CUDA 执行模型中的一层并发以及存储层次中的一个层级粒度逐级变细。线程块级 GEMMThreadblock-level GEMM每个线程块通过迭代地从全局内存加载输入矩阵的 tile并累加计算矩阵乘积从而算出输出 GEMM 中属于自己的那部分。在线程块层面数据从全局内存加载。分块策略总体上是达成高效的关键但程序员必须在多个相互冲突的目标之间权衡更大的线程块意味着更少的全局内存访存次数从而保证 DRAM 带宽不会成为瓶颈但过大的线程块 tile可能无法与问题维度良好匹配如果 GEMM 的M或N维度较小线程块中部分线程可能落在问题边界之外而做无意义的工作如果M和N都很小而K很大则该方案只会启动相对较少的线程块无法充分利用 GPU 内所有的多处理器SM。针对后一种情况文档 Parallelized Reductions 中描述的策略会沿 GEMM 的 K 维度把计算划分到多个线程块或多个线程束上各线程块/线程束并行计算矩阵乘积最后再归约得到最终结果。在 CUTLASS 中线程块 tile 的维度通过ThreadblockShape::{kM, kN, kK}指定可针对目标处理器与 GEMM 问题维度进行调优。从源码结构看该参数在 default_mma.h 等默认配置头中用于实例化线程块级MmaPipelined等主循环。线程束级 GEMMWarp-level GEMM线程束级 GEMM 映射到 CUDA 执行模型中的线程束级并行线程块内的多个线程束从共享内存把数据取入寄存器并执行计算。线程束级 GEMM 的实现方式有两种通过 Tensor Core 发射mma.sync/wmma指令通过发射到 CUDA 核心的线程级矩阵计算指令SIMT。为获得最高性能对共享内存的访问应做到无 bank 冲突bank conflict free为了最大化线程束内的数据复用应选择较大的线程束级 GEMM tile。线程级 GEMMThread-level GEMM在最低一级分块中每个线程负责处理一定数量的元素。线程之间无法访问彼此的寄存器因此需要选择一种组织方式使寄存器中的值可被多条数学指令复用。这导致线程内部形成一种 2D 平铺结构每个线程向 CUDA 核心发射一串相互独立的数学指令累加计算一个外积outer product。SGEMM、IGEMM、HGEMM、DGEMM均由线程级矩阵乘过程发射的 SIMT 数学指令计算得到。Epilogue上述代码只关注矩阵乘法C AB其结果保存在线程块内每个线程的寄存器中。输出 tile 的逻辑元素到线程的映射是为最大化矩阵乘性能而选择的但这种映射不会产生对全局内存高效的、合并coalesced的读写访问。Epilogue 是一个独立阶段线程通过共享内存交换数据然后以高效的条带化访问模式striped access patterns协作访问全局内存同时也是方便地利用矩阵乘结果作为输入、计算线性缩放linear scaling与其他逐元素运算的阶段。CUTLASS 定义了若干典型的 epilogue 运算如线性缩放与 clamp也允许使用其他设备端函数调用算子device-side function call operators执行自定义操作。优化策略Optimizations上述分层结构已经给出了到 CUDA 执行模型与 CUDA/Tensor Core 的高效映射。以下各节描述在各类问题规模角落获得峰值性能的策略在最大化并行度的同时尽可能利用数据局部性。软件流水线Pipelining分块结构要求每个 CUDA 线程的寄存器中容纳大量存储分配累加器accumulator元素通常至少占去线程寄存器预算的一半。因此相比于其他类型的 GPU 负载占用率occupancy即并发的线程/线程束/线程块数量相对较低这限制了 GPU 通过上下文切换到同一 SM 内其他并发线程来隐藏内存延迟与其他停顿的能力。为缓解内存延迟的影响CUTLASS 使用软件流水线software pipelining在线程内将内存访问与其他计算重叠。CUTLASS 通过以下两个作用域上的**双缓冲double buffering**实现线程块作用域的共享内存 tile在共享内存中分配两块 tile。一块用于加载当前矩阵运算所需数据另一块用于缓冲从全局内存加载的、供下一次 mainloop 迭代使用的数据线程束作用域的矩阵 fragment在寄存器中分配两份 fragment。一份在本次矩阵计算中传给 CUDA/Tensor Core另一份用于接收供下一次线程束级矩阵运算使用的共享内存取回结果。从源码结构看经典线程块级主循环 mma_pipelined.h 中的MmaPipelined继承自MmaBaseShape_, Policy_, 2并以static_assert((Base::kStages2), MmaPipelined requires kStages set to value 2)强制要求双缓冲流水线mma_base.h 中的kStages常量则控制共享内存 tile 的级数分配。下图展示了 CUTLASS GEMM 中使用的高效、流水化 mainloop 主体。线程块栅格化Threadblock Rasterization为最大化末级缓存last level cache中数据的复用CUTLASS 定义了若干函数来影响线程块到 GEMM 问题逻辑分区的映射把连续启动的线程块映射到被分区的 GEMM 问题的紧凑二维区域从而提高这些线程块在同一时刻访问同一批全局内存 tile 的概率。相关函数定义在 include/cutlass/gemm/threadblock/threadblock_swizzle.h 中例如GemmIdentityThreadblockSwizzle提供get_tiled_shape()把问题规模除以 tile 尺寸得到逻辑 tile 数、get_grid_shape()计算 CUDA 网格维度与get_log_tile()计算最优 swizzle 宽度等接口。在 gemm_splitk_parallel.h 等内核的Params构造中也会调用ThreadblockSwizzle().get_log_tile(grid_tiled_shape)来计算 swizzle 参数。并行化归约Parallelized ReductionsSplit K —— 跨线程块的归约矩阵乘积计算在O(MN)个相互独立的内积计算上暴露了并行性。对足够大的问题规模CUTLASS GEMM 内核可以逼近理论计算峰值但对小问题则没有足够的线程块来高效占满整个 GPU。作为对策将内积计算中执行的归约并行化可以在保持大线程块级 GEMM tile 吞吐优势的同时让更多线程块并发执行。CUTLASS 通过沿 GEMM 的 K 维度分区并为每个分区额外启动一组线程块来实现跨线程块的并行归约即 parallel reduction splitK 策略。该策略需要执行两个内核partitionedK GEMM类似于一种带 stride 的批处理batched stridedGEMM。它不需要用户指定每个 batch 的问题规模只需给定整体问题规模与沿 K 维应用于 A、B 操作数的分区数。例如参数 m128、n128、k4096、partition16将得到 16 个批处理 GEMM每批为 m128、n128、k256。partitionedK 也支持 K 无法被分区数整除的情形例如 m128、n128、k4096、partition20将得到 20 个批处理 GEMM前 19 批为 m128、n128、k4096/20204最后一批为 m128、n128、k220。batched reduction 内核以 partitionedK GEMM 的输出C为输入沿 K 维度执行归约。用户必须自行管理工作空间内存workspace memory来保存这一中间结果。从源码结构看partitionedK GEMM 由内核 gemm_splitk_parallel.h 中的GemmSplitKParallel实现其Params通过full_gemm_k_iterations / grid_tiled_shape.k()计算每个线程块的gemm_k_size并据此得出每个线程块沿 K 的偏移对最后一个分区threadblock_tile_offset.k() 1 grid_tiled_shape.k()则直接取问题剩余 K从而自然支持 K 不可整除的情形。Sliced K —— 跨线程束的归约与 split-k 类似sliced-k 旨在提升M、N 较小而 K 较大的内核效率。在线程块层面CtaTileN与CtaTileM通过把工作划分到线程束来暴露并行度更大的 warpTile 能带来更好的指令级并行ILP与复用但也会限制每个线程块内运行的线程束数量从而降低效率。为了提升这类场景的效率把 warpTile 也沿CtaTileK划分能让更多线程束在同一个 CTA 内并发运行从而更高效地利用硬件。Sliced-k 内核把线程块的计算不仅在CtaTileN、CtaTileM维度上划分到参与的线程束也在CtaTileK维度上划分。因此 sliced-k 会带来一点额外开销由于每个线程束只使用CtaTileK的一个切片slice归约前每个线程束只有部分和所以最后必须在参与的线程束之间执行一次归约。Hopper Warp 特化Warp Specialization注以下关于 warp 特化的内容针对 Hopper 内核设计。Blackwell SM100 内核具有截然不同的 warp 特化结构但将生产者producer与消费者consumer代理分离这一概念仍然适用。从 Hopper 起CUTLASS 3.0 将Warp Specialization概念纳入内核设计一个线程块被划分为两组线程束即生产者线程束组producer warp group与消费者线程束组consumer warp group。生产者线程束组使用新的Tensor Memory AcceleratorTMA把数据从全局内存加载到共享内存缓冲区。工作流程以 sm90_mma_tma_gmma_ss_warpspecialized.hpp 为参考实现生产者线程束组DMA等待消费者线程束组通过新增的Async Pipeline 类详见 pipeline.md将共享内存缓冲区信号置为empty数据写入共享内存后TMA 同时更新与该 stage 关联的 barrier通知相关线程缓冲区已filled消费者线程束组MMA等待生产者信号置为filled随后发射 Tensor Core MMA 运算最后消费者线程束组释放release缓冲区供下一轮 TMA 加载使用。从源码结构看内核 sm90_gemm_tma_warpspecialized.hpp 中通过WarpGroupRole::Producer / Consumer明确区分线程束组角色并为不同角色配置MainloopPipeline与EpiLoadPipeline的ThreadCategory。Warp 特化持久协作内核设计Warp-Specialized Persistent Cooperative从 Hopper 开始引入的另一种 Warp 特化内核设计是 sm90_gemm_tma_warpspecialized_cooperative.hpp 中的Warp-Specialized Persistent Cooperative内核。与 Warp 特化内核相同线程束组的概念与线程束组之间的 barrier 同步保持不变。其独特之处在于持久线程块persistent thread blocks按 kernel_hardware_info.hpp 中KernelHardwareInfo结构体给出的 SM 数量占满 GPU这些持久线程块用于平铺输出从而在其生命周期内可能计算多个输出 tile。主要收益是摊薄了所有内核都有的线程块启动与内核前导prologue开销存在两个消费者线程束组协作计算同一个输出 tile把 tile 沿 M 维度一分为二。这允许启用更大的 tile 尺寸——因为每个消费者线程束组的寄存器压力降低了——从而提升性能。由于每个线程块现在计算多个输出 tile网格启动的形状与 tile 到线程块的调度由新的Tile Scheduler管理见 sm90_tile_scheduler.hpp其PersistentTileSchedulerSm90继承自StaticPersistentTileScheduler。Tile Scheduler 会考虑cluster 的形状以及可用的SM 数量计算出输出 tile 到已启动线程块的合法调度。Warp 特化持久 Ping-Pong 内核设计Warp-Specialized Persistent Ping-Pong第三种内核设计是 sm90_gemm_tma_warpspecialized_pingpong.hpp 中的Warp-Specialized Persistent Ping-Pong内核。与 Persistent Cooperative 内核相同线程束组概念、线程束组间 barrier 同步与网格启动形状保持一致。其独特之处在于两个消费者线程束组通过 Tile Scheduler 被分配到不同的输出 tile从而让一个消费者线程束组的epilogue与另一个消费者线程束组的数学运算重叠——最大化 Tensor Core 利用率生产者线程束组使用Ordered Sequence Barrier见 sm90_pipeline.hpp 中OrderedSequenceBarrier类其构造模板参数SequenceDepth与SequenceLength控制序列深度与长度同步按顺序依次填充两个消费者线程束组的缓冲区。从源码结构看Ping-Pong 内核通过WarpGroupRole::Consumer0 / Consumer1区分两个消费者线程束组sm90_gemm_tma_warpspecialized_pingpong.hpp并使用LoadWarpOrderBarrier cutlass::OrderedSequenceBarrier1,2与MathWarpGroupOrderBarrier实现有序填充。进一步阅读以下仓库内资源提供了 GEMM 设计与实现的更多细节pipeline.mdAsync Pipeline 与 barrier 同步机制详解gemm_api.md 与 gemm_api_3x.mdCUTLASS 2.x 与 3.x 的 GEMM API 使用说明code_organization.md源码目录组织与模块划分cutlass_3x_design.mdCUTLASS 3.x 内核设计总览quickstart.md 与 getting_started.rst构建、编译与快速上手。Copyright本文内容整理自 CUTLASS 仓库 media/docs/cpp/efficient_gemm.md版权归 NVIDIA CORPORATION AFFILIATES2017 - 2026所有遵循 SPDX-License-Identifier: BSD-3-Clause 许可协议发布。【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门