ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

CUTLASS Blackwell FMHA 示例完全解析:前向、生成、反向与 MLA 的实现、构建与运行指南

CUTLASS Blackwell FMHA 示例完全解析:前向、生成、反向与 MLA 的实现、构建与运行指南 CUTLASS Blackwell FMHA 示例完全解析前向、生成、反向与 MLA 的实现、构建与运行指南【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass本指南以 examples/77_blackwell_fmha/README.md 为核心骨架逐层展开 CUTLASS 面向 NVIDIA BlackwellSM100/SM103架构的融合多头注意力FMHA示例覆盖前向context/generation 两阶段、反向含 MLA backward与 MLA 推理weight-absorbed 权重吸收三种场景。读者可以从中掌握各场景支持的数据类型与 HeadDim 约束、blocking 与加载方式TMA 与cp.async的选择依据、变长序列的正确使用姿势padding 与total_length以及如何通过collective/fmha_fusion.hpp自定义融合mask / activation的扩展点同时结合源码给出 CMake 构建、测试参数与命令行选项的完整解读。示例概览与目录结构examples/77_blackwell_fmha/是 CUTLASS 仓库中用于演示 Blackwell 架构 FMHA 的完整示例集合共包含 6 个 CUDA 源文件与 12 个构建产物fp8/fp16 各半。其核心思想README 明确指出并在源码中得到印证是复用 collective GEMM builder 的算子选择逻辑把选型结果重新组合成 FMHA kernel再针对 FMHA 定制 kernel 层与 collective 层。目录结构如下examples/77_blackwell_fmha/ ├── 77_blackwell_fmha.cu # 前向context 阶段 ├── 77_blackwell_fmha_gen.cu # 前向generation 阶段KV cache ├── 77_blackwell_mla.cu # MLA 推理2SM 大 latent head ├── 77_blackwell_mla_fwd.cu # MLA 前向 ├── 77_blackwell_fmha_bwd.cu # FMHA 反向含 MLA backward ├── 77_blackwell_fmha_2bwd.cu # FMHA 反向两遍 GEMM 版本 ├── collective/ # 自定义 collectivemainloop / epilogue / 加载 ├── device/ # 设备级入口fmha.hpp / sm100_mla.hpp ├── kernel/ # FMHA 专用 kernel、tile scheduler、选项机制 ├── reference/ # 参考实现与误差校验 ├── CMakeLists.txt └── README.md从源码结构看示例的关键设计是kernel 与 collective 层全部被表述为“fmha-specific”的形式——collective/下的 mainloop、epilogue、加载模块都以sm100_fmha_*命名kernel/下则包含前向 kernelsm100_fmha_fwd_kernel_tma_warpspecialized.hpp、生成阶段 kernelsm100_fmha_gen_kernel_warpspecialized.hpp、反向 kernelsm100_fmha_bwd_kernel_tma_warpspecialized.hpp以及 MLA 专用 kernelsm100_fmha_mla_tma_warpspecialized.hpp 等。FMHA for Blackwell前向Forward / Context / Generation支持范围与 blocking 建议前向示例覆盖HeadDim 为 32、64、128输入数据类型为 fp8、fp16、bf16。针对两种使用场景README 给出了明确的 blocking 建议场景M-blockingSeqlen-QN-blockingSeqlen-K加载方式Forward / Context256128TMAGeneration128Num-Groups实际上限目前为 3264 / 128 / 256cp.asyncContext 阶段通过 TMA 完成加载而 generation 阶段使用cp.asyncREADME 解释其原因是cp.async更适合复杂的加载模式如 KV cache 中带 remap、cache-only 等场景。生成阶段每个 threadblock 处理一个 group因此 M 维上的“Num-Groups”即并发组数README 特别注明虽然 blocking 设为 128但实际的 Num-Groups 上限目前为 32。在源码层面这一设计由加载模块分工实现TMA 路径在 sm100_fmha_load_tma_warpspecialized.hppcp.async路径在 sm100_fmha_load_cpasync_warpspecialized.hpp。两者配合不同的 mainloop 使用sm100_fmha_fwd_mainloop_tma_warpspecialized.hpp 对应 TMA 前向sm100_fmha_gen_mainloop_warpspecialized.hpp 对应生成阶段。核心设计两个 tile 的 pingpongREADME 描述的算法结构是每个 threadblock 分配两个 tile在矩阵乘法与 softmax 之间对这两个 tile 做 pingpong交替使用。这一结构直接对应软硬件流水线当一个 tile 的 QK GEMM 产生累加器后紧接着对其做 softmax含 mask 应用随后送入 PV GEMM两个 tile 轮转以隐藏延迟。在 sm100_fmha_fwd_mainloop_tma_warpspecialized.hpp 中可以看到实现细节QK 与 PV 两个 GEMM 各自通过cutlass::gemm::collective::CollectiveBuilder以Sm100OpClassTensorOpKernelTmaWarpSpecialized1SmSm100选型构建第 90-105 行ThreadShape注释说明两个 softmax warp 的排布(2, 1, 1)表示堆叠排布对 Q 大、加载 K/V 最少的场景最优(1, 2, 1)表示并排排布对小 Q / 大 K 最优Q 的 stage 数为 2KV 的 stage 数为sizeof(Element_) 1 ? 4 : 3fp8 元素占 1 字节时用 4 级否则 3 级K 与 V 的共享内存通过union复用第 115-118 行。QK GEMM 的累加器类型为 fp32ElementAccumulatorQK floatPV 阶段同样使用 fp32 累加ElementAccumulatorPV float这保证了 softmax 与最终输出累加的数值稳定性输出类型在 fp8 构建时为cutlass::float_e4m3_t否则为cutlass::half_t见 77_blackwell_fmha.cu。同时输出 scale 会传入 collective mainloop在量化前以 FP32 compute 完成缩放源码文件头注释“Output Scale”一节。变长序列Variable Sequence Length的正确用法这是 README 中反复强调、且带版本演进的要点4.3.0 变更务必按以下规则使用padding 内存要求代码要求第一个输出 batch 之前存在一批“有效但从不使用”的 padding 内存valid but never used padding memory ahead of the first output batch输入张量无需 padding但要求输入张量不含 NaN 或 Inf 值必须把total_length设置为problem_shape这是容易踩坑的细节若忘记设置kernel 拿到的 Q/K 总长度会与实际分配的内存不一致。源码佐证VariableLength结构体collective/fmha_fusion.hpp有三个成员——max_length最大序列长度、cumulative_length累积长度指针、total_length默认 -1。在 TMA 加载路径中sm100_fmha_load_tma_warpspecialized.hpp当cumulative_length_q与cumulative_length_k都非空时QK problem shape 会直接用total_length覆盖第 0、1 维反向 kernel 中同样以total_length作为 Q、K 的总长度参与地址计算见 sm100_fmha_bwd_dkdv_kernel_tma_warpspecialized.hpp。而示例主程序初始化 varlen 问题时正是用VariableLength{max_seqlen_q, nullptr, total_seqlen_q}构造 problem shape77_blackwell_fmha.cu其中第三个参数即total_length。VariableLength还配套了一组模板工具fmha_fusion.hppapply_variable_length把 shape 中的VariableLength叶子替换为cumulative_length[idx1] - cumulative_length[idx]并生成带偏移的坐标apply_variable_length_offset则同时返回调整后的 shape 与基于cumulative_length[idx]的偏移量——这些工具被 tile scheduler 与加载逻辑用来把每个 batch 的局部坐标映射到拼接后的全局内存。融合扩展点collective/fmha_fusion.hppREADME 指出collective/fmha_fusion.hpp是最容易的融合定制点。核心机制是每个 mask 类型NoMask、ResidualMask、CausalMaskkIsQBegin、CausalForBackwardMask等都实现get_trip_count/get_masked_trip_count/get_unmasked_trip_count/apply_mask四个接口apply_mask接收第一个 GEMMQK的累加器以及这些元素对应的逻辑位置IndexQK非常适合实现 mask 或 activation如把元素置为-INFINITY以在 softmax 前屏蔽若融合需要内存加载比如读取 bias、alibi 表等则不能只改apply_mask需要修改 mainloop collective通过 TMA 编排这些加载。以CausalMask的实现为例fmha_fusion.hppget_trip_count计算该线程块实际需要迭代的 K tile 数。IsQBegintrueQ 位于矩阵开头默认时取min(max_blocks_k, ceil_div((blk_q1)*TileM, TileN))IsQBeginfalseQ 位于末尾通常用于只计算下一行的推理场景时还要加上offset_q seqlen_k - seqlen_q的偏移get_masked_trip_count/get_unmasked_trip_count把迭代拆成“需要 mask 的 tile”与“无需 mask 的 tile”两部分前者可以走完整的掩码路径后者可跳过掩码计算apply_mask逐元素判断(get0(pos) get1(pos)) || (get1(pos) seqlen_k)并置为-INFINITY。ResidualMask则用于seqlen_k % TileN ! 0的剩余元素屏蔽这是 TMA 无法透明处理的情况注释明确指出d % kHeadDim ! 0或seqlen_q % kBlockM ! 0会被 TMA 与 epilogue 的 predication 透明处理无需额外 mask。命令行参数Forward前向示例的完整参数见 77_blackwell_fmha.cu 的print_usage参数含义默认值 / 备注--bintbatch 数 B未指定时由16384/k推导源码parse逻辑--hintQ 的头数 H未指定时h 2048 / d--h_kintK/V 的头数GQA/MQA未指定时等于 h--qintQ 序列长度未指定时取 k--kintK/V 序列长度未指定时取 q--varlen-qint:int:...每个 batch 的 Q 变长冒号分隔与--q互斥--varlen-kint:int:...每个 batch 的 K 变长冒号分隔与--k互斥--dinthead dimension128--tensor_ring_buffersinttensor ring 缓冲区数1--warmup_iterationsintwarmup 迭代数1--iterationsint计时迭代数3--verify与参考实现比对验证默认关闭--verbose打印 smem 与每个 kernel 的执行时间默认关闭--maskno\|residual\|causal掩码类型空/no表示无掩码--causal-typeqbegin\|qend因果掩码类型默认qbegin--persistent启用 persistent scheduler默认关闭--varlen启用变长序列BQ 与 BK 变为总长度并按 batch 拆分--sm-count指定 SM 数而非查询—--kernel-filterfilter按正则匹配 kernel—其中--init-style[-q/-k/-v]支持r(random)/1(ones)/d(linear stride 1)/s(linear stride 128)/n(none) 五种初始化方式。注意parse中的约束--varlen与--q/--k不能同时给出77_blackwell_fmha.cuvarlen 模式下--b必须与--varlen-q/--varlen-k的元素个数一致。一个基本用法示例源码文件头注释给出./examples/77_blackwell_fmha/77_blackwell_fmha_fp8 --b2048 --h2048 --d2048 --q2048 --k2048FMHA for Blackwell生成阶段fmha_gen77_blackwell_fmha_gen.cu对应 README 中 generation 阶段的前向面向KV cache 推理。CMake 测试展示了其独有的能力CMakeLists.txt--remap启用 batch index remappingcache 中不同 batch 的 KV 可以按任意顺序摆放kernel 通过索引重映射找到对应条目--cache-only只使用 KV cache 中已有数据不读取也不插入新条目--varlencache 条目间的序列长度可变--clear-cache计时运行前清空 cache。注意 README 的 Changes 一节明确指出fmha_gen示例仅支持 head dim 1284.1.0 澄清因此 CMake 中TEST_GEN_HDIM64被注释掉。CMake 中TEST_GEN_BASIC使用--b1 --h4 --k512 --d128 --verifyTEST_GEN_GQA/TEST_GEN_REMAP/TEST_GEN_CACHEONLY则组合了--h_k2与 remap/cache-only 开关CMakeLists.txt。该路径的 mainloop 与加载模块分别是 sm100_fmha_gen_mainloop_warpspecialized.hpp 与 sm100_fmha_load_cpasync_warpspecialized.hpp印证 README 所述 generation 使用cp.async。FMHA for Blackwell反向Backward反向示例支持HeadDim 为 64 与 128数据类型 fp8、fp16、bf16Q 与 K 的 blocking 均为 128加载走 TMA并支持 causal maskingREADME 注明反向目前不支持 GQACMake 中相关测试注释为“bwd doesnt support GQA yet, --h_k will just get ignored”。反向计算由三个 kernel 组成README 列出注意源码中实际顺序与编号FmhaKernelBwdSumOdO源码 fmha_kernel_bwd_sum_OdO.hpp计算 O 与 dO 外积之和Sm100FmhaBwdKernelTmaWarpSpecialized源码 sm100_fmha_bwd_kernel_tma_warpspecialized.hpp计算反向主体dQ、dK、dVREADME 强调它是本示例的核心看点——展示如何用 tensor core 构建高性能融合 kernelFmhaKernelBwdConvert源码 fmha_kernel_bwd_convert.hpp把 dQ 从 fp32 转换到最终输出精度。从 kernel 命名可以推断反向主体在源码层面还拆分为 dKdV kernelsm100_fmha_bwd_dkdv_kernel_tma_warpspecialized.hpp与 dQ kernelsm100_fmha_bwd_dq_kernel_tma_warpspecialized.hpp二者共同构成Sm100FmhaBwdKernelTmaWarpSpecialized的完整流水。反向的 mask 类型使用CausalForBackwardMaskfmha_fusion.hpp它在CausalMask与ResidualMaskForBackward基础上合并判断(pos_q offset_q pos_k) || !elem_less(pos, problem_size)时置-INFINITY。MLA Backwardd192, d_vo128README 说明反向示例还提供了MLAMulti-head Latent Attentionbackward能力--d192 --d_vo128。其核心 kernel 是Sm100FmhaBwdMlaKernelTmaWarpSpecialized源码 sm100_fmha_bwd_mla_kernel_tma_warpspecialized.hppREADME 强调 MLA 路径对原始算法做了针对性调整以在 MLA 形状下获得高性能。CMake 中的验证用例为TEST_BWD_MLA_BASIC--b1 --h4 --q512 --k512 --d192 --d_vo128 --verify --maskno与TEST_BWD_MLA_VARLEN变长 --maskresidual。MLA Inference for Blackwell权重吸收推理背景与数据范围MLA 推理示例在weight-absorbed权重吸收机制下工作latent head dim 512、rope head dim 64支持 fp16、bf16、fp8 输入与输出。由于 latent head 维度很大输出累加器accumulator规模也随之增大README 说明该示例通过使用 2 个 SM 的 Blackwell tensor core2Sm来承载这一形状——这也是生成物命名为77_blackwell_mla_2sm_*的原因。加载方式TMA 或 cp.async paging加载路径有三条README 归纳加载方式适用限制TMA无 paging无TMApage size 128支持 paging 的 TMA 路径cp.async支持任意 ≤128 的 2 的幂 page size启用 paging 后代码还支持变长序列。对应的实现文件为 sm100_fmha_mla_load_tma_warpspecialized.hppTMA与 sm100_fmha_mla_fwd_mainloop_tma_warpspecialized.hppMLA 前向 mainlooppipeline 辅助在 common/pipeline_mla.hpp2 的幂工具在 common/pow_2.hpp。六个二进制与三种模式README 说明该示例构建6 个二进制对应 fp8/fp16 × 三种模式TMA 模式77_blackwell_mla_2sm_{fp8,fp16}无 CPASYNC 宏cp.async 模式77_blackwell_mla_2sm_cpasync_{fp8,fp16}编译时定义CPASYNC宏见 CMakeLists.txtback-to-back GEMMB2B模式把 softmax 变成 no-op用于隔离/剖析 GEMM 本身。同一份77_blackwell_mla.cu通过CPASYNC编译宏区分 TMA 与cp.async两条路径。CMake 为 MLA 目标附加了-Xptxas -v以便在编译时输出寄存器/共享内存占用信息CMakeLists.txt。MLA 与 split-KV、reduction 策略CMake 测试还揭示了 MLA 在长序列下的两种 reduction 策略CMakeLists.txtTEST_MLA_SEP_REDUCTION--b1 --k4096 --split_kv8 --page128 --verify——分离式 reduction先各自算部分 softmax再统一合并代码见 sm100_fmha_mla_reduction.hppTEST_MLA_FUSE_REDUCTION同样形状但加--fuse_reduction——融合式 reductionTEST_MLA_LARGE_SPLIT_KV--verify --split_kv20 --page128验证更大 split 因子。MLA 的 tile 调度由 sm100_mla_tile_scheduler.hpp 负责设备入口在 device/sm100_mla.hpp。构建与运行CMake 集成构建条件与产物CMakeLists.txt 的关键约束是示例只在非 Windows、非 Clang 编译器、且CUTLASS_NVCC_ARCHS匹配100a或103a即 SM100/SM103Blackwell / Blackwell Ultra时才参与构建第 116 行。README 的 Changes 一节确认4.4.0 起支持 Blackwell UltraSm103。构建产物的完整清单77_blackwell_fmha_all目标77_blackwell_fmha_fp8 / fp16 # 前向 context 77_blackwell_fmha_gen_fp8 / fp16 # 前向 generation 77_blackwell_mla_2sm_fp8 / fp16 # MLA TMA 77_blackwell_mla_2sm_cpasync_fp8 / fp16 # MLA cp.async 77_blackwell_fmha_bwd_fp8 / fp16 # 反向 77_blackwell_mla_fwd_fp8 / fp16 # MLA 前向另有77_blackwell_fmha_2bwd_fp8/fp16在 foreach 循环内构建。所有示例源文件都带--use_fast_math -ftemplate-backtrace-limit0编译标志第 39 行前者启用快速数学后者限制模板回溯输出长度以便定位模板实例化错误。测试即文档CMake 中的参数组合README 建议通过CMakeLists.txt中的测试或每个二进制的--help来了解调用方式。CMake 中的测试参数正是最完整的“用法文档”可归纳为几组基础与掩码前向与反向共用--b1 --h4 --q512 --k512 --d128 --verify --maskno --b1 --h4 --q512 --k512 --d128 --verify --maskcausal --verify --iterations0 --b1 --h1 --h_k1 --q1013 --k1024 --d128 --maskcausal --causal-typeqend --b1 --h4 --q512 --k512 --d128 --verify --maskresidual --varlen --b2 --h4 --q512 --k512 --d64 --verify # HeadDim64 --b2 --h4 --h_k2 --q512 --k512 --d64 --verify # GQA变长序列矩阵TEST_VARLEN_00至TEST_VARLEN_22前向与 MLA 前向各一套覆盖--varlen-q/--varlen-k的不同组合包括常规变长--d128 --h8 --h_k4 --varlen-q128 --varlen-k128不同 batch 不同长度--varlen-q256:256:256:256 --varlen-k256:768:512:512零长度 batch--varlen-k256:0:1280:512、--varlen-q256:0:512:256验证对空序列 batch 的处理极小长度--varlen-q3:2 --varlen-k2:5、--varlen-q17:10 --varlen-k13:10、--varlen-q1 --varlen-k1超长非整除--varlen-q177:845 --varlen-k257:766causal 变长--causal-typeqbegin与--causal-typeqend分别搭配--varlen-q17 --varlen-k257等形状以及q1013/k1024、q1024/k1035这类 Q 与 K 不相等的情形。生成阶段--b1 --h4 --k512 --d128 --verify # basic --b1 --h4 --k512 --d128 --verify --varlen # 变长 cache --b2 --h4 --h_k2 --k512 --d128 --verify # GQA --b2 --h4 --h_k2 --k512 --d128 --verify --remap # 索引重映射 --b2 --h4 --h_k2 --k512 --d128 --verify --cache-onlyMLA--b1 --k512 --page128 --verify # basic paging --b1 --k4096 --split_kv8 --page128 --verify # 分离 reduction --b1 --k4096 --split_kv8 --page128 --fuse_reduction --verify --verify --split_kv20 --page128 # 大 split 因子MLA 反向--b1 --h4 --q512 --k512 --d192 --d_vo128 --verify --masknoMLA backward 开关即--d192 --d_vo128。验证机制所有测试都带--verify把 CUTLASS kernel 的输出与 reference/fmha_fwd_reference.hpp前向参考、reference/fmha_bwd_reference.hpp反向参考、reference/fmha_mla_reference.hppMLA 参考比对。前向示例的 verify 逻辑77_blackwell_fmha.cu同时检查 O 与 LSElog-sum-exp两个输出fp8sizeof(Element)1时最大差阈值 1e-1、平均差阈值 1e-1其他精度最大差 1e-2、平均差 1e-3。GQA/MQA 的 stride 机制在源码中有清晰注释77_blackwell_fmha.cuhead dimension 表示成元组K/V 在第一维 stride 为零即退化为 MQAQ 用(grouped_heads, heads_kv)、K/V 用(grouped_heads, heads_kv):(0, head_stride)即 GQA。版本演进与兼容性提示README 的 Changes 一节给出该示例随 CUTLASS 版本的关键演进4.1.0增强变长序列测试禁用 MLA 的 B2B 模式以简化示例澄清fmha_gen示例仅支持 head dim 1284.3.0变长序列新增“输出前需要 padding batch”与“输入不含 NaN/Inf、必须设置total_length为problem_shape”的约束前文已详述4.4.0新增 Blackwell UltraSm103支持。结合构建条件CUTLASS_NVCC_ARCHS匹配100a或103a可知本示例面向 Blackwell 家族专用指令集TCGEN05 tensor core不适用于旧架构若目标机器为 Ampere/Hopper应使用 CUTLASS 仓库中其他 FMHA 示例如examples/fmha系列。小结examples/77_blackwell_fmha完整展示了在 CUTLASS 3 中为 Blackwell 构建融合注意力 kernel 的工程范式以 collective GEMM builder 完成算子选型再以 FMHA 语义重写 kernel 与 collective。三个方向前向/生成、反向、MLA共享同一套掩码抽象fmha_fusion.hpp的apply_mask与 trip-count 机制、同一套变长序列抽象VariableLengthtotal_length与同一套 tile 调度思想但针对各自场景选择不同加载原语TMA vscp.async与不同 blocking。无论是想为自有模型接入 Blackwell FMHA还是希望学习如何在 CUTLASS collective 层做算子融合这个示例都是直接的参考起点——建议从collective/fmha_fusion.hpp的apply_mask入手它同时是 README 指定的最小融合扩展点。参考文件索引示例说明examples/77_blackwell_fmha/README.md构建与测试参数examples/77_blackwell_fmha/CMakeLists.txt前向主程序examples/77_blackwell_fmha/77_blackwell_fmha.cu生成阶段主程序examples/77_blackwell_fmha/77_blackwell_fmha_gen.cu融合扩展点mask/activationexamples/77_blackwell_fmha/collective/fmha_fusion.hpp前向 TMA mainloopexamples/77_blackwell_fmha/collective/sm100_fmha_fwd_mainloop_tma_warpspecialized.hpp反向 kernelexamples/77_blackwell_fmha/kernel/sm100_fmha_bwd_kernel_tma_warpspecialized.hppMLA 反向 kernelexamples/77_blackwell_fmha/kernel/sm100_fmha_bwd_mla_kernel_tma_warpspecialized.hppMLA 前向 kernelexamples/77_blackwell_fmha/kernel/sm100_fmha_mla_tma_warpspecialized.hpp参考实现examples/77_blackwell_fmha/reference/fmha_fwd_reference.hpp【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表