
CUTLASS Grouped Kernel Schedulers 深度解析Problem Visitor 调度原理、Rank2K 优化与负载均衡实践【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlassCUTLASSCUDA Templates and Python DSLs for High-Performance Linear Algebra的 grouped kernel 是典型的持久化persistentkernel它把多个问题如 GEMM、SYR2K、HER2K放进同一次 CUDA kernel 启动中执行。本文以官方文档 grouped_scheduler.md 为主体结合仓库源码如 grouped_problem_visitor.h 与 rank_2k_grouped_problem_visitor.h系统讲解 grouped kernel scheduler代码中称为problem visitor如何为每个线程块分配 tile、两种调度模式kDeviceOnly与kHostPrecompute的实现差异以及如何通过按 K 维排序问题来改善负载均衡。读完后你将理解 grouped kernel 的核心调度机制并能结合源码为你的多问题场景选择正确的调度模式。CUTLASS Gemm Kernel 层级结构图一、Grouped Kernel Scheduler 是什么CUTLASS 的 grouped kernel 是一种持久化 kernel它在一次 CUDA kernel 启动内处理多个问题例如多个 GEMM、多个 SYR2K/HER2K。与普通 CUTLASS GEMM 不同——后者启动的线程块数量等于 GEMM 中 tile 的数量——grouped kernel 通常启动的线程块数量少于所有问题 tile 的总数。每个线程块负责计算组内一个或多个问题上的若干 tile。grouped kernel 的scheduler代码中被称为problem visitor负责为每个线程块分配它需要计算的 tile 序列。其核心思想是线程块持久运行一个循环不断向调度器询问下一个要计算的 tile 并执行该 tile 对应的 kernel 级操作MMA 与 epilogue。伪代码如下ProblemVisitor problem_visitor; while (problem_visitor.next_tile()) { // // Get next tile index from scheduler // // // Compute MMA and epilogue // // Inform the scheduler that we are done with the current tile problem_visitor.advance(gridDim.x); }调度器的关键功能集中在next_tile()方法中它决定了调用线程块接下来应计算组内哪个 tile如果还有可计算 tile 的话。对应源码中BaseGroupedProblemVisitor::advance()的实现为tile_idx grid_size;即每次完成一个 tile 后按网格规模步进详见 grouped_problem_visitor.h。二、Grouped GEMM Schedulerround-robin 分配grouped GEMM 使用的调度器以round-robin轮询方式将组内的 tile 分配给线程块。例如考虑一组包含 4 个 GEMM、每个 GEMM 都是 2x2 tile 网格的情形假设启动 8 个线程块。下图展示了每个 GEMM 中每个 tile 被分配到的线程块 IDALT当问题拥有不同数量的 tile 时映射关系如下ALT2.1 计算某个 block 的调度grouped GEMM 中的每个线程块通过调用上文所述的next_tile()方法自行计算调度。为此线程块的ProblemVisitor维护一个tile_idx成员初始化为blockIdx.x每计算完一个 tile 就递增gridDim.xgrouped kernel 的启动配置只使用 x 维度。调度器随后需要确定tile_idx属于组中的哪个 GEMM以及它映射到该问题中的哪个 tile确定tile_idx属于哪个 GEMM调度器从最近访问过的 GEMM 开始向后遍历将该 GEMM 内的 tile 数量累加到一个运行变量problem_tile_start上。当满足problem_tile_start tile_idx problem_tile_start tiles_in_problem时就找到了正确的 GEMM。确定tile_idx对应 GEMM 内的哪个 tile找到 GEMM 后该 block 要计算的 tile 由tile_idx - problem_tile_start给出。随后执行简单的光栅化rasterization把这个一维 tile ID 映射到 GEMM 中的二维 tile 坐标。源码层面BaseGroupedProblemVisitor中tile_idx、problem_tile_start、problem_idx这三个成员与文档描述完全对应见 grouped_problem_visitor.h构造时以tile_idx(block_idx), problem_tile_start(0), problem_idx(0)初始化L121-L123。在 Scheduler Modes 一节中我们描述这种搜索是如何被加速的。三、Grouped Rank2K Scheduler为三角矩阵特化上一节描述了 grouped GEMM kernel 所用调度器的工作原理。虽然该调度器足以正确实现 grouped Rank2K 操作即 SYR2K 和 HER2K但它会带来显著的效率问题。3.1 grouped GEMM 调度器用于 grouped Rank2K 的缺陷grouped GEMM 调度器假设组内每个 GEMM 的每个 tile 最终都会影响问题的输出。但 Rank2K 问题并非如此其矩阵 C 是上三角或下三角。对这类问题使用默认的 grouped GEMM 调度器会导致线程块频繁被分配到会提前退出early exit的 tile例如被分配到下三角问题中位于上三角区域的 tile。这进一步造成线程块之间的负载不均衡因为 grouped GEMM 调度器给所有线程块分配了几乎相同的 tile 数量而不管其中真正活跃的 tile 有多少。考虑一个包含 4 个 SYR2K 问题的例子每个问题的矩阵 C 由 2x2 的 tile 网格组成矩阵 C 是下三角阴影 tile 表示。假设启动 8 个线程块计算这一组问题。默认 grouped GEMM 调度器会按以下顺序分配线程块ALT在这个例子中线程块 1 和 5 被持续分配到非活跃 tile。当组内问题尺寸各不相同时我们观察到这种方案依然会导致显著的负载不均衡。3.2 为三角问题特化调度器我们希望能设计一个调度器对于输出矩阵为三角的 kernel能更高效地把线程块映射到活跃 tile。理想情况下调度器只把线程块分配给下三角问题的下三角区域内的 tile上三角问题则相反。基于上面的例子这样的调度器产生的线程块到 tile 的分配可能如下ALT实现该调度需要从线程块 ID 映射到 tile 坐标(i, j)。下面以 3x3 网格的下三角矩阵为例说明映射方法。我们首先按从 1 开始的行、列、tile 与线程块 ID 计算行列索引再减 1 转换为 0 起始的版本。该映射方法在很大程度上借鉴了 Stack Overflow 上的描述https://stackoverflow.com/a/40954159。ALT由线程块 IDt计算行i对于给定的行 i该行内所有线程块 ID t 都满足t 1 2 3 ... (i-1) i右侧的闭式公式为i(i1)/2。据此可由 t 解出 it i(i1)/2 2t i^2 i 2t i^2 i 0.25 - 0.25 2t 0.25 i^2 i 0.25 2t 0.25 (i 0.5)^2 sqrt(2t 0.25) - 0.5 i为了处理小数部分取i ceil(sqrt(2t 0.25) - 0.5)转换为 0 起始的行并处理 0 起始的 ti ceil(sqrt(2(t1) 0.25) - 0.5) - 1 ceil(sqrt(2t 2.25) - 0.5) - 1由线程块 IDt与行i计算列j对于给定的行 i该行内所有线程块 ID t 还满足t 1 2 3 ... (i-2) (i-1) -- t i(i-1)/2一行内的线程块 ID 是连续的因此对于 1 起始的线程块 ID t 和行 i1 起始的列 ID 为j t - (i(i-1)/2)其 0 起始版本为j (t1) - (i(i1)/2) -1 t - (i(i1)/2)处理非方形网格尽管 Rank2K 问题的整体输出尺寸保证是方形的但由于可能使用非方形的线程块形状实际计算所用的网格可能不是方形的。例如线程块形状 64x32 作用在 128x128 的输出问题上会得到 2x4 的 tile 网格。这个情形可以这样处理注意到输出形似一个 2x2 的“宏 tile”macro tile方形网格每个宏 tile 内包含 2 个“真 tile”。因此可以先利用上面的公式把线程块 ID 映射到其“宏 tile”再映射到宏 tile 内的“真 tile”。以 2x4 网格为例映射过程如下ALT0 起始的线程块 IDt到“宏 tile ID”t_macro的映射为t_macro t // r其中r是网格最大维度与最小维度的比值上例中r 4 / 2 2。用t_macro和上面的公式计算方形矩阵中的行与列得到i_macro和j_macro0 起始。从(i_macro, j_macro)到(i, j)的映射非常简单if (ThreadblockShape::M ThreadblockShape::N): r ThreadblockShape::M / ThreadblockShape::N i i_macro j (j_macro * r) (t % r) elif (ThreadblockShape::M ThreadblockShape::N): r ThreadblockShape::N / ThreadblockShape::M i (i_macro * r) (t % r) j j_macro else: i i_macro j j_macro这段逻辑在源码中有直接对应Rank2KGroupedProblemVisitor::threadblock_offset()中先计算macro_row与macro_col对kUpper模式交换二者再通过OffsetHelper::macro_row_to_row / macro_col_to_col展开到真实 tile 坐标见 rank_2k_grouped_problem_visitor.h。其中行计算ceil(sqrt((2*macro_id) 2.25) - 0.5) - 1正是上文推导的闭式公式的直接实现。网格维度互不为倍数的情况即使线程块形状 M 和 N 通常是彼此的倍数某个问题的网格维度也可能与线程块的比值不一致。例如132x132 的问题使用 64x32 的线程块形状会得到 3x5 的 tile 网格此时每个“宏 tile”内没有整数个“真 tile”。遇到这种情况时只需把网格较大的维度补齐pad使得每个“宏 tile”内有整数个“真 tile”。于是上面例子中的 3x5 网格会被当作 3x6 网格处理。每个 tile 的行列位置照常计算。凡是映射到问题范围之外或上/下三角区域之外的 tile例如 (2, 5)的线程块都会从该问题提前退出并可能继续处理组中的下一个问题。处理上三角矩阵对于上三角矩阵唯一的改动是在上述计算中交换i_macro与j_macro。源码中Rank2KGroupedProblemVisitor通过if (kFillModeC cutlass::FillMode::kUpper) { cutlass::swap(macro_row, macro_col); }实现这一点并静态断言只允许kLower或kUpper填充模式rank_2k_grouped_problem_visitor.h。源码佐证Rank2K 的 tile 计数在 Rank2KGroupedProblemSizeHelper::tile_count() 中CUTLASS 只统计对角线及其下方或kUpper时上方的 tile 数tiles_on_diagonal dimtiles_below_diagonal dim * (dim - 1) / 2再乘以OffsetHelper::kThreadblockSkewRatio处理长宽维度不同的情况。这从实现层面印证了 Rank2K 调度器只向活跃 tile 分配线程块的设计目标。四、Scheduler Modes两种调度模式grouped kernel 调度器提供了两种查找下一个 tile 的模式由枚举cutlass::gemm::kernel::GroupScheduleMode控制其定义位于 grouped_problem_visitor.h/// Enumerated type describing the type of scheduling to perform for the ProblemVisitor enum class GroupScheduleMode { // Perform all scheduling on device kDeviceOnly, // Precompute on the host the full sequence of problems to access kHostPrecompute };4.1GroupScheduleMode::kDeviceOnly默认该模式在设备端完成所有调度工作。它通过让 warp 内每个线程“拥有”一个不同的问题并判断tile_idx是否落在这个问题的范围内从而并行化对tile_idx所属问题的搜索。kDeviceOnly以 warp 级warp-wide方式并行warp 内每个线程按其 lane id 加载一个问题尺寸并计算该问题的 tile 数量随后用 warp 级前缀和prefix sum找出该 warp 所考察问题的起始 tile。前缀和结束时每个线程持有组中一个唯一问题的起始 tile 索引与 tile 计数。只要tile_idx仍处于 warp 当前托管问题的范围内每个线程都会检查tile_idx是否落于自己当前问题的范围。匹配的问题索引及其起始 tile 随后被广播broadcast给 warp 内所有线程。源码中的实现与上述描述一致GroupedProblemVisitor..., kDeviceOnly, ...中每个线程按 lane 计算problem_ending_tile随后用__shfl_up_sync循环做 warp 内包含式前缀和最后通过__ballot_sync与__popc确定tile_idx所在的问题索引见 grouped_problem_visitor.h。4.2 在主机端预计算调度GroupScheduleMode::kHostPrecompute该模式通过在主机端预计算每个 block 将访问的问题序列减少设备端执行的调度量。如前所述要把tile_idx映射到某个问题内的具体 tile只需问题 ID 和该问题的起始 tile相对组内所有 tile 而言。因此该调度器为每个 block 计算的每个 tile 预计算问题索引与问题起始 tile。单个 block 的调度表示为一个(problem_idx, problem_starting_tile)元组数组每个 block 对应一个数组。这些数组在主机端生成并拷贝到设备端。这种表示针对“每个 block 在每个问题上至多计算一个 tile”的情形做了优化。当一个 block 在组内同一个问题上计算多个 tile 时上述表示会产生重复条目因此是次优的例如一个 block 在问题 3 上计算两个 tile而问题 3 的起始 tile 索引是 20就会得到[(3, 20), (3, 20)]。CUTLASS 之所以选择这种表示是因为 grouped kernel 本身在问题尺寸较小时通常收益最大此时 block 在每个问题上至多计算一个 tile。从源码看kHostPrecompute模式通过get_workspace_size()计算工作区大小sizeof(ProblemInfo) * entries_per_block * block_count在host_precompute()中把每个 tile 映射为(problem_idx, problem_start)写入主机端工作区并拷贝到设备端运行时next_tile()直接从共享内存中预取的prefetched_problems数组读取下一项见 grouped_problem_visitor.h。kHostPrecompute还要求PrefetchTileCount 0即必须配合共享内存预取使用L344-L345。4.3 该选哪种调度模式选择调度模式时请考虑以下问题grouped kernel 的输入参数如 ptrA、lda在你的应用中如何设置如果这些参数由设备上先前运行的 kernel 设置而非由主机设置你可能希望使用kDeviceOnly因为它能把额外的主机-设备通信降到最低。你的应用中主机端工作能否与其他设备 kernel 重叠例如若 grouped GEMM 被用作神经网络中的第 N 层grouped GEMM 的主机端预计算可以潜在与第 N-1 层的设备端工作重叠。这种情况下kHostPrecompute可能更合适。组内问题的计算强度如何kHostPrecompute与kDeviceOnly的性能差异在计算强度低的 grouped kernel 上最为明显因为此时调度器耗时占 grouped kernel 运行时间的比重较大。直观地说随着组内问题计算强度的下降MMA 操作消耗的运行时间占比减小调度逻辑消耗的时间占比增大。由于两种调度模式只影响 grouped kernel 的调度逻辑因此计算强度较低的组更可能从kHostPrecompute中获益。五、通过排序问题改善负载均衡grouped kernel 调度器会给参与 kernel 的每个 block 分配几乎等量的 tile。组内每个 tile 的 M、N 维度相同但每个 tile 的 K 维度取决于所属问题的 K 维度因此不同 tile 的 K 维度可能不同。而 tile 的 K 维度在很大程度上决定了该 tile 的计算耗时。5.1 K 维度不均衡的潜在问题要保证各 block 之间的计算负载均衡重要的一点是每个 block 计算的所有 tile 的 K 维度之和应当与其他 block 相近。如果某个 block 计算的大 K tile 远多于其他 block它可能比其他 block 耗时更长。例如考虑下面一组 GEMM0 1152x768x128 1 1152x768x1024 2 768x1152x128 3 768x1152x1024若 tile 尺寸为 128x128则每个问题有 54 个 tile组内共 216 个 tile。假设该 grouped GEMM 运行在拥有 108 个 SM 的 GA100 上且其占用率occupancy为 1——即每个 SM 同时只能有一个活跃线程块。于是 grouped GEMM 将运行 108 个持久线程块每个 block 计算(216 / 108) 2个 tile。在 grouped GEMM 调度器采用的 round-robin tile 分配下本组 GEMM 的 tile 分配如下Threadblocks 0-53: Tiles of size 128x128x128 from problem 0 Threadblocks 54-107: Tiles of size 128x128x1024 from problem 1 Threadblocks 0-53: Tiles of size 128x128x128 from problem 2 Threadblocks 54-107: Tiles of size 128x128x1024 from problem 3按照这一分配线程块 54-107 的工作量明显大于线程块 0-53前者计算两个 K1024 的 tile后者只计算两个 K128 的 tile。由于分配不均衡线程块 54-107 的运行时间会显著长于线程块 0-53导致线程块 0-53 在大量时间内空闲。显然对本例更好的 tile 分配方式是让所有线程块各计算一个 K1024 的 tile 和一个 K128 的 tile这样能更好地均衡各 block 的工作量。5.2 通过排序降低不均衡一个简单且可能降低负载不均衡的办法是把组内问题按K 维度降序排序。这之所以能改善负载均衡是因为组内 tile 是按 round-robin 方式顺序分配给 block 的因此每个 block 总是会被分配当前可用 K 维度最大的下一个 tile。回到上面的例子在执行 grouped GEMM 前先排序问题尺寸在 GA100 上两种调度模式下该 grouped GEMM 的运行时间都改善了约 30%。为了简化“按此方式排序问题及其关联元数据”的过程设备级 grouped kernel 提供了sort_problems()方法。其源码实现位于 base_grouped.h它用std::stable_sort按problem_sizes_ptr[i].k()降序生成索引然后用reorder_array同步重排问题尺寸以及lda/ldb/ldc/ldd与 A/B/C/D 的偏移量指针。实际用法可以参考 grouped GEMM 示例 examples/24_gemm_grouped。最后需要提醒虽然排序问题在某些场景下有效但它并不保证一定提升性能。某些情况下由于影响 GEMM 性能的其他冲突因素排序后性能反而可能下降。我们建议对你的 grouped kernel 分别在有排序和无排序的情况下做 profiling确认排序在你的场景中是否有帮助。六、总结调度核心grouped kernel 是持久化 kernel每个线程块通过ProblemVisitor::next_tile()循环领取任务kDeviceOnly在设备端用 warp 级并行与前缀和完成问题查找kHostPrecompute在主机端预生成(problem_idx, problem_start)序列并预取到共享内存。Rank2K 特化由于输出为三角矩阵grouped GEMM 的 round-robin 调度会浪费线程块CUTLASS 通过“宏 tile”闭式映射含非方形网格补齐、上三角交换行列只向活跃三角区域分配线程块对应实现见 rank_2k_grouped_problem_visitor.h。负载均衡tile 的 K 维度决定其耗时按 K 降序排序问题sort_problems()可显著改善均衡性文档示例中 GA100 上提升约 30%但需以实际 profiling 结果为准。选型建议参数由设备端 kernel 生成时优先kDeviceOnly主机端预计算可与前序设备工作重叠、且组内计算强度较低时优先尝试kHostPrecompute。版权声明本文内容基于 CUTLASS 开源仓库文档与源码整理。Copyright (c) 2017 - 2026 NVIDIA CORPORATION AFFILIATES. All rights reserved. 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),仅供参考