资讯动态

CUTLASS Hopper 混合精度 Grouped GEMM 实战:基于 Example 69 的 Mixed-input Grouped GEMM 实现与调优

发布时间:2026/9/15 16:07:13 来源:尧图企业网站定制
CUTLASS Hopper 混合精度 Grouped GEMM 实战基于 Example 69 的 Mixed-input Grouped GEMM 实现与调优【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass导读本文围绕 CUTLASS 仓库中的示例 examples/69_hopper_mixed_dtype_grouped_gemm/README.md深入讲解如何在 NVIDIA Hopper 架构SM 90上使用 CUTLASS 3.x 新一代 Collective API 实现混合输入类型Mixed-input的 Grouped GEMM即在一个批次内同时执行多个问题形状M、N、K各不相同的 GEMM且矩阵 A 与 B 使用不同的数据类型如 BF16×E5M2、FP8×INT4。读完本文你将掌握 Grouped GEMM 与普通混合精度 GEMM 在接口上的差异布局指针类型、分组参数与 stride 数组的传入方式、示例程序的核心模板配置、命令行参数的实际用法、性能测试方法以及当前实现的能力边界。1. 从 Example 55 到 Example 69为什么需要 Grouped 混合精度 GEMMExample 69 是 examples/55_hopper_mixed_dtype_gemm 的扩展Example 55 展示了单个 GEMM 中 A、B 具有不同类型Mixed-input时的 CUTLASS 3 API 用法Example 69 在其基础上引入Grouped GEMM能力——将多个问题形状可能不同的 GEMM 打包到同一次 kernel 启动中批量执行这一模式在 MoEMixture of Experts推理、动态形状推理等场景中非常常见。从源码注释69_hopper_mixed_dtype_grouped_gemm.cu可以看到示例的接口与 Example 55 的标准混合输入 GEMM 高度相似核心差异可以归纳为三点这也是 README 明确给出的三条要点Collective Builder 中把布局类型替换为布局指针类型例如把LayoutB_Transpose换成LayoutB_Transpose *表示每个 group 的矩阵可以拥有独立的指针指向各自的数据从而支持不同 group 使用不同基址。Arguments 中传入分组信息包括 group 数量groups、问题尺寸数组每个 group 的M, N, K以及矩阵 A、B 各自的 stride 数组。若包含 scale / zero-point还需要额外传入它们的 stride 数组下文StrideS相关代码即是佐证。此外README 特别提醒了一个容易混淆的参数命名问题Example 55 中--g表示量化 scale 的组大小group size of scaling而 Example 69 中--groups表示GEMM 的个数。为避免语义冲突Example 69 改用--c表示 scale 的 chunk/组大小。读者在使用两个示例时务必区分这三个参数。2. 示例目录结构与三个可执行目标examples/69_hopper_mixed_dtype_grouped_gemm/目录共包含 6 个文件文件作用69_hopper_mixed_dtype_grouped_gemm.cu主示例BF16A× E5M2B混合输入 Grouped GEMM69_hopper_int4_fp8_grouped_gemm.cuINT4B× FP8A示例带离线 layout swizzle 与 LUT/PRMT 转换69_hopper_int4_bf16_grouped_gemm.cuINT4B× BF16A示例grouped_mixed_dtype_utils.hpp三个示例共享的命令行解析、问题生成、性能统计工具CMakeLists.txt三个可执行目标的注册与测试命令定义README.md本文所依据的说明文档其中CMakeLists.txt通过cutlass_example_add_executable将三个.cu文件分别注册为69_hopper_mixed_dtype_gemm、69_hopper_int4_fp8_grouped_gemm、69_hopper_int4_bf16_grouped_gemm并挂载了十余组测试参数TEST_RANDOM、TEST_FIXED、TEST_SMALL、TEST_SCALE_PERCOL、TEST_SCALE_GROUP等这些测试命令可视为官方对该实现正确性验证的参考用例。3. 构建与运行3.1 环境前提源码在 69_hopper_mixed_dtype_grouped_gemm.cu 的main()中做了两道硬性检查CUDA Toolkit 版本必须 ≥ 12.3__CUDACC_VER_MAJOR__ 12 || (__CUDACC_VER_MAJOR__ 12 __CUDACC_VER_MINOR__ 3)时报错提示。GPU 必须是 NVIDIA Hopper 架构compute capability 9.0props.major ! 9 || props.minor ! 0时提示 This example requires a GPU of NVIDIAs Hopper Architecture。同时kernel 相关代码整体包在CUTLASS_ARCH_MMA_MODIFIABLE_TMA_SM90_SUPPORTED宏内说明它依赖 SM90 上可修改 TMA 描述符的硬件特性。因此该示例只面向 Hopper GPU不适用于 Ampere 及更早架构。3.2 运行命令README 和源码头部给出的标准运行命令为$ ./examples/69_hopper_mixed_dtype_grouped_gemm/69_hopper_mixed_dtype_grouped_gemm --m2048 --n2048 --k2048 --mode1 --groups10行为说明来自源码注释 69_hopper_mixed_dtype_grouped_gemm.cu上例中 10 个 group 都使用给定的--m/--n/--k作为统一尺寸如果省略某个维度参数如只给--m该维度会在各 group 之间随机化--alpha/--beta同理若省略则每个 group 的 alpha/beta 会被随机赋值见initialize()中(rand() % 5) 1与rand() % 5的生成逻辑。3.3 全部命令行参数grouped_mixed_dtype_utils.hpp的print_usage()完整列出了参数表参数类型默认值含义--help标志—显示用法说明--mint整数1024所有 group 的 M 维度--nint整数2048所有 group 的 N 维度--kint整数512所有 group 的 K 维度--cint整数512量化权重 scale 的 chunk 大小等价于 Example 55 的--g--groupsint整数6Grouped GEMM 中独立 GEMM 问题的个数--modeint整数1运行模式见下--alphaf32浮点1.0Epilogue 标量 alpha--betaf32浮点0.0Epilogue 标量 beta--iterationsint整数100性能统计迭代次数--warmupint整数10预热迭代次数--benchmarkstr字符串空从文件加载 benchmark 问题形状其中--mode的取值沿用 mixed_dtype_utils.hpp 中的枚举0ConvertOnly直接执行A BB 不缩放对应主示例中的GemmConvertOnly1ScaleOnly执行A (scale * B)对应GemmScaleOnly2ScaleWithZeroPoint执行A (scale * B zero-point)——注意该模式当前在 Grouped 版本中并未实现main()只分派 mode 0 和 1 两个分支见 69_hopper_mixed_dtype_grouped_gemm.cu。3.4 问题形状的生成规则与 benchmark 文件格式GroupedMixedDtypeOptions::parse()会根据--benchmark是否为空决定问题形状来源未指定--benchmark调用randomize_problems()。对每个 group若命令行给出了某维度则固定使用否则随机生成——其中 M 随机为alignment * ((rand() % 64) 1)N、K 使用默认值。若 K 不是 alignment 的整数倍会抛出runtime_errork dimension must be a multiple of ...。这里的 alignment 定义为tma_alignment_bits / sizeof_bitsQuantType即 128 bit TMA 对齐约束换算成的元素个数例如 E5M2 为 128/816。指定了--benchmarkpath按文本文件逐行读取idx MxNxK如0 2048x5120x8192每个维度会向上取整到 alignment 的整数倍groups自动取文件中的问题个数。3.5 alpha / beta 的两种传参方式在args_from_options()69_hopper_mixed_dtype_grouped_gemm.cu中epilogue 融合参数根据命令行是否提供 alpha/beta 分为两种形态统一标量命令行给出了--alpha和--beta则所有 group 共用同一标量alpha_ptr/beta_ptr/alpha_ptr_array/beta_ptr_array均置空dAlpha/dBeta步长为{0, 0, 0}。逐 group 数组未给出时传入alpha_ptr_array/beta_ptr_array指向设备端每个 group 的 alpha/beta 数组dAlpha/dBeta步长变为{0, 0, 1}即按 group 索引步进。这是 Grouped GEMM 区别于普通 GEMM 的关键点之一每个 group 可以拥有独立的 alpha、beta 与问题形状。4. 核心代码深度解析4.1 问题形状GroupProblemShape三个示例统一使用using ProblemShape cutlass::gemm::GroupProblemShapeShapeint,int,int; // M,N,K per group其定义位于 include/cutlass/gemm/group_array_problem_shape.hpp持有num_groups、设备端problem_shapes指针与主机端host_problem_shapes指针并提供groups()、get_problem_shape(group_idx)、get_host_problem_shape(group_idx)、is_host_problem_shape_available()四个接口。同文件还定义了适用于 MoE 的MoEProblemShape以tokens_per_expert描述每个 expert 的 token 数与统一单问题的ArrayProblemShape可以看出 CUTLASS 3 以统一的“问题形状抽象”来覆盖单 GEMM、Ptr-Array GEMM、Grouped GEMM 与 MoE 四类场景。4.2 布局指针类型Grouped 化的关键编译期改造以主示例BF16×E5M2为例Collective 配置如下69_hopper_mixed_dtype_grouped_gemm.cuusing CollectiveMainloopConvertOnly typename cutlass::gemm::collective::CollectiveBuilder ArchTag, OperatorClass, ElementB, LayoutB_Transpose *, AlignmentB, ElementA, LayoutA_Transpose *, AlignmentA, ElementAccumulator, TileShape, ClusterShape, cutlass::gemm::collective::StageCountAutoCarveout static_castint(sizeof(typename CollectiveEpilogue::SharedStorage)), KernelSchedule ::CollectiveOp;与普通 GEMM 相比关键变化正是 README 强调的把布局类型替换为布局指针类型LayoutB_Transpose *、LayoutA_Transpose *epilogue 侧同理使用LayoutTransposeLayoutC::type *与LayoutTransposeLayoutD::type *。指针类型的存在意味着 mainloop 将通过“指针数组”Ptr-Array方式访问每个 group 的独立数据缓冲。对应的调度策略为KernelSchedule cutlass::gemm::KernelPtrArrayTmaWarpSpecializedCooperativeEpilogueSchedule cutlass::epilogue::PtrArrayTmaWarpSpecializedCooperativeTileShape 为Shape_128,_16,TileShapeK其中TileShapeK 128 * 8 / sizeof_bitsMmaTypeClusterShape 为Shape_1,_1,_1从调度策略命名可以看出该实现基于 Hopper 的 TMATensor Memory Accelerator与 Warp-Specialized Cooperative 流水线stage 数通过StageCountAutoCarveout在扣除 epilogue 共享内存占用后自动取最大。4.3 Scale 的配对方式tuple 打包当 B 需要缩放时源码注释明确指出“Scale 信息必须与被缩放的 operand 配对”因此using CollectiveMainloopScaleOnly typename cutlass::gemm::collective::CollectiveBuilder ArchTag, OperatorClass, cute::tupleElementB, ElementScale, LayoutB_Transpose *, AlignmentB, ElementA, LayoutA_Transpose *, AlignmentA, ...ElementB与ElementScale被包进cute::tuple作为 B 侧操作数的“元素类型”表示 B 矩阵与 scale 向量一同流入 mainloop 并在 kernel 内部完成反量化。在 INT4×FP8 示例中scale 甚至打包为cutlass::ArrayElementScale, 8见 69_hopper_int4_fp8_grouped_gemm.cu并在--shuffle开关下使用离线 layout swizzleLayoutAtomQuanttile_to_shape优化 INT4 数据在全局内存中的读取顺序。4.4 Arguments 组装Grouped 模式下的数据传递args_from_options()中按ConversionMode分派两类 arguments69_hopper_mixed_dtype_grouped_gemm.cu// DirectConvertmode 0 arguments { cutlass::gemm::GemmUniversalMode::kGrouped, {options.groups, problem_sizes.get(), nullptr}, // group 数 设备端问题尺寸数组 {ptr_B.get(), stride_B.get(), ptr_A.get(), stride_A.get()}, // B、A 的指针数组与 stride 数组 {fusion_args, ptr_C.get(), stride_C.get(), ptr_D.get(), stride_D.get()}, hw_info }; // ConvertAndScalemode 1额外追加 scale 指针数组、scale stride 数组与 chunk 大小 c arguments { cutlass::gemm::GemmUniversalMode::kGrouped, {options.groups, problem_sizes.get(), nullptr}, {ptr_B.get(), stride_B.get(), ptr_A.get(), stride_A.get(), ptr_scale.get(), stride_S.get(), options.c}, {fusion_args, ptr_C.get(), stride_C.get(), ptr_D.get(), stride_D.get()}, hw_info };可见 README 所述“传入 group 数量、问题尺寸数组、A/B 的 stride 数组以及 scale 的 stride 数组”在代码中一一对应。hw_info通过query_device_multiprocessor_count获取 SM 数量用于 kernel 内部的资源规划所有分配集中在allocate()中完成——为每个 group 累计M*K、K*N、M*N等元素数后一次性分配整块显存再通过offset_*数组切片出各 group 的指针。4.5 正确性验证逐 group 参考 GEMMverify()69_hopper_mixed_dtype_grouped_gemm.cu针对每个 group 执行参考流程先用cutlass::dequantize()在 kernel 外完成B → B_dq的反量化即把“kernel 内缩放”与“kernel 外缩放”两条路径对齐到同一参考语义用标准非 Grouped的GemmRefGemmUniversalKernelScheduleAuto对每个 group 独立计算参考输出用BlockCompareRelativelyEqual以相对误差epsilon1e-2、非零下限1e-4与 CUTLASS kernel 输出逐组比对并打印Group i: M,N,K, alpha, beta, Status。run()中还先执行一次正确性/预热 run再进入grouped_mixed_dtype_profiling()做性能统计。4.6 性能统计口径grouped_mixed_dtype_utils.hpp中的grouped_mixed_dtype_profiling()与gflops()实现了与 Example 55 一致的计时方案用cudaEvent在warmup iterations次gemm.run()循环中计时仅统计iter warmup的样本平均耗时avg_runtime_ms取所有有效样本均值吞吐按gflops 2 * Σ(M_i × N_i × K_i) / (runtime_s × 1e9)计算——将所有 group 的 FLOP 累加后折算每个乘加计 2 次浮点操作。程序最终输出Groups、Avg runtime (ms)、GFLOPS以及Disposition: Passed/Failed。5. 官方测试矩阵与参数参考CMakeLists.txt中定义的测试组合是理解参数语义的现成参考注意测试默认--iterations0以关闭性能统计、只做正确性检查set(TEST_RANDOM --iterations0) # 随机问题形状 set(TEST_RANDOM_LARGE_GROUP --groups100 --iterations0) set(TEST_EPILOGUE --alpha0.5 --beta0.5 --iterations0) set(TEST_EPILOGUE_LARGE_GROUP --alpha2.0 --beta2.0 --groups100 --iterations0) set(TEST_FIXED --m2048 --n5120 --k8192 --groups16 --iterations0) # 固定问题形状 set(TEST_SMALL --m256 --n128 --iterations0) set(TEST_RANDOM_PERF --iterations10) # 性能测试 set(TEST_DIRECT_BATCHED --m2048 --n5120 --k8192 --mode0 --iterations0) # 直接转换 set(TEST_SCALE_PERCOL --m4096 --n5120 --k8192 --c8192 --mode1 --iterations0) # 每列缩放 set(TEST_SCALE_GROUP --m2048 --n5120 --k8192 --c512 --mode1 --iterations0) # 分组缩放其中两个 scale 测试特别值得注意--c8192时scale_k ceil_div(k, 8192)当k ≤ 8192时每组仅有 1 个 scale等价于逐列per-column缩放--c512表示每 512 个 K 元素共享一个 scale即分组group-wise缩放。这正对应 README “Upcoming features” 一节中的能力边界当前 Mixed-input Grouped GEMM只支持 row-wise scalingscale 维度为scale_k × N见allocate()中elements_scale scale_k * N并且group-wise scaling 要求所有 group 的问题形状相同zero-point 与 block-wise scaling 尚未支持源码中block_zero虽有分配与初始化但 mode 2 分支未接入分派逻辑args_from_options()的else分支对非法 mode 直接报错退出。6. 与 Example 55 的对照速查维度Example 55普通混合精度 GEMMExample 69Grouped 混合精度 GEMM问题形状单个M,N,KGroupProblemShapeShapeint,int,int每 group 一个M,N,KCollective 布局参数LayoutA/LayoutBLayoutA*/LayoutB*指针类型scale 组大小参数--g--c避免与--groups混淆GEMM 个数参数--lbatch 数--groupsalpha/beta全局标量支持逐 group 指针数组alpha_ptr_array反量化时机kernel 内kernel 内可选或参考路径 kernel 外dequantize7. 局限性总结与选型建议综合 README 与源码当前实现截至仓库当前版本的边界如下只支持 row-wise scaling即 scale 向量沿 K 方向分块chunk size 由--c控制group-wise scaling 仅适用于所有 group 问题形状相同的场景否则缩放语义无法统一对齐不支持 zero-pointScaleWithZeroPoint模式未实现与block-wise scaling运行前提是 CUDA Toolkit ≥ 12.3 与 HopperSM 90GPUkernel 依赖 SM90 的 Modifiable TMA 特性CUTLASS_ARCH_MMA_MODIFIABLE_TMA_SM90_SUPPORTED。因此如果你的场景是 MoE/动态批量推理、且 A/B 异构、每个 expert 的形状可能不同且量化方案为行方向/分组 scale无 zero-pointExample 69 的接口GroupProblemShape 指针布局 GemmUniversalMode::kGrouped即是可直接参照的 CUTLASS 3 标准写法若需要 zero-point 或 block-wise scaling则需等待后续版本扩展或参考 Example 55 的单 GEMM 路径自行组合。版权说明示例与文档版权归 NVIDIA CORPORATION AFFILIATES 所有Copyright (c) 2017 - 2026采用 BSD-3-Clause 许可SPDX-License-Identifier: BSD-3-Clause完整许可证文本见 README.md 与各源文件头部。【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价