Mojo 中的 Multi-Head Flash Attention 设计:从在线 Softmax 到 Hopper 上的 Warp 特化 FA3 内核 📅 发布时间:2026/9/12 16:42:18 👁 浏览次数: Mojo 中的 Multi-Head Flash Attention 设计从在线 Softmax 到 Hopper 上的 Warp 特化 FA3 内核【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo导读本文以 multi-head-flash-attention.md 为核心系统讲解 Modular PlatformMAX Mojo仓库中多Head注意力MHA如何基于 Flash Attention 算法实现从自注意力的数学定义出发推导 Flash Attention 2 的在线 softmax 分块算法再深入到针对 Hoppersm90架构特化的 Flash Attention 3FA3——包括wgmma异步矩阵乘、共享内存屏障、动态寄存器分配与 warp 组特化ping pong 内核。读完本文你将掌握 FA2/FA3 的核心算法推导、数值稳定性处理rowmax/rowsum 校正并能对照仓库中的真实实现max/kernels/src/nn/attention/gpu/nvidia/sm90/mha.mojo与基准配置理解一个生产级 MHA 内核的调度、分块与掩码设计。背景自注意力与多Head注意力自注意力机制self-attention的定义为A softmax(QK)V其中Q、K、V分别是该注意力头attention head对应的查询queries、键keys和值values。多Head注意力在此基础上增加参数q_heads与kv_heads约束为q_heads % kv_heads 0。令group q_heads // kv_heads则A softmax(Q_{q_head} K_{kv_head}) V_{kv_head} softmax(Q_{q_head} K_{q_head//group}) V_{q_head//group}这样数组就只用q_head索引即可描述。再引入批次维度batch_idx实际要执行的操作是softmax(Q_{batch_idx, q_head_idx} K_{batch_idx,q_head_idx//group}) V_{batch_idx,q_head_idx//group}这意味着本质上需要执行batch_size * num_q_head次注意力计算但只需加载batch_size * num_kv_head份唯一的K和V矩阵——这正是 GQAGrouped Query Attention减少 KV 访存量的来源。将Q、K、V视为 4 维ragged数组Q: batch_size x seq_len x num_q_heads x depth K: batch_size x num_keys x kv_num_heads x depth V: batch_size x num_keys x kv_num_heads x depth按batch_idx与q_head_idx索引后每次注意力计算包含三步S Q K P softmax(S) O P V用seq_len x depth的Q乘上depth x num_keys的K对seq_len x num_keys的结果矩阵做逐行 softmax用seq_len x num_keys的P乘上num_keys x depth的V。朴素算法的数据移动瓶颈朴素算法的主要成本在于数据移动。depth通常很小例如128而seq_len与num_keys可能很大——文档中以 llama3.3.70b 为例seq_len可达8192、num_keys可达119132。若直接物化一个8192 x 119132的中间矩阵将产生极高的内存带宽开销而对depth128这 128 个元素的归约根本无法掩盖这样的访存代价。Flash Attention 的核心创新正是避免物化这个大型中间矩阵把输出保持在寄存器中并采用在线 softmaxonline softmax进行计算。Flash Attention 2在线 Softmax 与分块算法数值稳定性为什么要减 rowmaxsoftmax 的朴素形式为softmax(S[i,j]) exp(S[i,j]) / sum(exp(S[i,k]) for k in range(num_keys)) exp(S[i,j]-S[i,a]) / sum(exp(S[i,k]-S[i,a]) for k in range(num_keys))其中a是行内最大值rowwise maximum。对于 32 位浮点数当x 88.72284时exp(x)Inf。为防溢出并保持数值精度减去行最大值可保证最大的指数项为1.0。在线算法的推导在线算法允许把整行计算拆成多个批次。先不应用分母直到最后才作用到最终输出数组上因此只需关注分子的更新以及此前 tile 计算出的输出。设b为先前批次的旧最大值索引exp(S[i,j]-S[i,a]) exp(S[i,j]-S[i,b]S[i,b]-S[i,a]) exp(S[i,j]-S[i,b])*exp(S[i,b]-S[i,a])要更新旧值只需乘以校正因子exp(S[i,b]-S[i,a])。这要求在整个在线算法中持续跟踪rowmax值以及最终作为分母的指数和rowsum。FA2 算法伪代码有了上述推导就可以按列分块K。Flash Attention 2 算法本质如下文档原文代码row_max [-Inf for _ in range(seq_len)] row_sum [0 for _ in range(seq_len)] O matrix(seq_len, depth).fill(0) for kv_start in range(0, num_keys, BN): block_range range(kv_start, kv_startBN) S mask_function(Q K[:, block_range]) # apply mask, e.g. CausalMask() old_rowmax rowmax row_max max(old_rowmax, rowmax(S)) P exp(S - row_max) correction exp(old_rowmax - row_max) row_sum row_sum * correction rowsum(P) O correction*O P V[block_range, :] O / row_sum注意各行之间没有通信因此很自然地按seq_len分块。这样所有临时量row_max、row_sum、S、P、O的规模都有界可以选取合适尺寸让它们全部驻留在寄存器中。由此避免了物化大矩阵以及随之而来的读写开销唯一的写操作是最终的输出唯一的读操作是输入Q、K、V。仓库中的在线 softmax 实现印证在仓库的 softmax.mojo 中可以看到与上述推导一一对应的底层原语_rowmax_online_softmaxmax/kernels/src/nn/softmax.mojo#L2695计算当前分数 tile 的行最大值并更新rowmax_tensor_rowsummax/kernels/src/nn/softmax.mojo#L2819对exp后的 tile 做行求和_online_softmax_correctionmax/kernels/src/nn/softmax.mojo#L2898当新 tile 带来更大行最大值时用exp(old_rowmax - new_rowmax)校正此前累积的 rowsum 与输出——这正是文档中correction exp(old_rowmax - row_max)的实现。这些函数被max/kernels/src/nn/attention/gpu/nvidia/sm90/mha.mojo直接 import 使用构成了从算法推导到生产内核的完整链路。特殊场景Token 生成与 KV-Cache一个重要的特例是 token 生成解码。利用 KV-cache 可以保存先前的结果后续计算使用seq_len1逐步增量地产生新结果。此时用kv_head_idx索引数组每次处理group行即共享同一 KV head 的一组 query head算法其余部分类似。进一步的优化流水线与缓冲进一步的优化包括使用缓冲进行流水线化例如用异步的 global → shared memory 拷贝。这有助于隐藏延迟当计算当前迭代时同时提前拷贝num_pipeline_stages - 1个迭代所需的数据。仓库的MHAConfig见 mha_utils.mojo将num_pipeline_stages默认设为4FA3 路径下 K/V 的共享内存按num_pipeline_stages * block_n * padded_depth分配kv_smem_size(fa3True)正是该思想的工程化体现。Flash Attention 3Hopper 架构上的 Warp 组特化FA3 是专门为 Hoppersm90架构优化的 Flash Attention 版本。Hopper 带来了三大新特性异步wgmma指令用于异步执行矩阵乘共享内存屏障shared-memory barriers支持跨 warp 组的细粒度同步动态寄存器分配/释放支持 warp 组特化warp-specialization。Ping PongWarp 组特化内核Warp 组特化内核也常被称为 ping pong 内核其思想是让不同 warp 组各司其职一个 warp 组专精于发起到共享内存的异步拷贝并释放自身绝大部分寄存器其余的 warp 组执行计算。GPU 的执行可以在它们之间来回弹跳ping pong同时这些操作在后台异步运行。通过运行两个计算 warp 组可以让它们的矩阵乘与 softmax 相关指令如指数运算相互重叠——因为这两类操作在硬件上使用不同的执行单元能够完全并行。上下文解码场景单 Warp 组内的流水线在 warp 组内部进一步优化这也是上下文解码的唯一选项因为此时Q的行数不足以划分成两个 warp 组可以对前面的循环做流水线化文档原文代码for kv_start in range(BN, num_keys, BN): # copy from P (f32 register tile) to a bf16 tile # this frees the f32 register tile for S Q K[...] S Q K[:, range(kv_start, kv_startBN)] O P V[range(kv_start-BN, kv_start), :] # from the previous iteration! S.wait() S mask_function(S) # apply mask, e.g. CausalMask() old_rowmax rowmax row_max max(old_rowmax, rowmax(S)) P exp(S - row_max) correction exp(old_rowmax - row_max) row_sum row_sum * correction rowsum(P) O.wait() # frees bf16 register tile O correction*O # NOTE: the P in P V is a truncation to bfloat16, so the registers # do not alias S or P elsewhere;注意其中P V[subset, :]的P是截断truncation为bfloat16的副本因此寄存器不会与别处的S、P产生别名冲突。这样P V的wgmma指令可以重叠内核内绝大部分向量指令而所有这些操作又都能与另一个 warp 组负责的内存传输相重叠从而大规模并行利用硬件资源。仓库中的 FA3 内核印证仓库在 max/kernels/src/nn/attention/gpu/nvidia/sm90/mha.mojo 中实现了完整的 SM90Hopper H100FA3 MHA 内核与分发层模块 docstring 明确说明这是warp-specialized, TMA-based MHA kernel for NVIDIA H100 GPUs。从源码中可以观察到的关键设计硬性约束检查comptime assert BM % 64 0、BK % 64 0H100 使用 128B swizzle、algorithm FlashAttentionAlgorithm(3)并强制config.dtype KVType.dtype q_typemha.mojo#L189-L229。生产者线程FA3 内核为生产者保留1 * WARPGROUP_SIZE即 128个线程num_threads num_consumer_threads 128对应 mha_utils.mojo 中num_producer_threads[producer_consumer_kernel]()返回 128 的逻辑。TMA 与 swizzleQ/K/V 分别通过QTMATile、KVTMATile构建 TMA tile 描述符使用TensorMapSwizzle.SWIZZLE_128BK、V 共享内存 tile 分别采用 k-major / mn-major 布局mha.mojo#L1014-L1031以匹配wgmma的访存要求。调度器通过TransientScheduler/TileScheduler/QueuedTileScheduler把scheduler、sink、KV row offsets、ragged valid lengths 等comptime 配置物化为具体内核实例USE_EXPERIMENTAL_KERNELS宏可开启 persistent kernel 实验路径。掩码mask_function对应的CausalMask实现在 mha_mask.mojo#L461 附近注释明确确保 token 只受先前 token 影响此外还提供ChunkedCausalMask、SlidingWindowCausalMask、RelativeLogitsMask等变体供不同场景选择。FlashAttentionAlgorithm结构体mha_utils.mojo#L111定义了四个变体NAIVE、FLASH_ATTENTION_1、FLASH_ATTENTION_2、FLASH_ATTENTION_3默认构造值为FLASH_ATTENTION_3当值未指定-1时init()会根据目标 dtype 与 GPU 架构自动选择最优算法——sm90/sm100 且为半精度或 sm100 且 fp8时选 FA3否则回退到 FA2。实践验证仓库中的基准与形状配置仓库提供了完整的 FA MHA 基准配置 bench_flash_attention.yaml其中以 llama3 常见形状为公共参数llama3-commons: llama3-commons mask_rank: 4 qkv_type: DType.bfloat16 mask_type: DType.bfloat16 depth: 128 num_heads: 32 group: 4 $batch_size: [1] $mode: flash_attention即depth128、num_heads32、group4对应kv_heads 8、batch size 为 1并针对seq_len/num_keys从32一路扫到16384的多种序列长度做基准。这与文档中seq_len和num_keys可以很大llama3.3.70b 中分别可达8192与119132的动机描述相互印证FA 正是为了让这种长序列注意力不再受S×S中间矩阵的物化开销所困。总结从本文可以看到一条清晰的演进主线数学层多Head注意力把问题归结为batch_size * num_q_head次softmax(QK)V计算GQA 让 KV 只需加载kv_head份算法层FA2 用在线 softmaxrowmax/rowsum 跟踪与指数校正把S×S的物化矩阵消除所有中间量驻留寄存器仅读写Q/K/V与最终输出token 生成借助 KV-cache 退化为seq_len1的增量计算架构层FA3 针对 Hopper 引入wgmma异步矩阵乘、共享内存屏障与动态寄存器分配通过 warp 组特化ping pong与单 warp 组内流水线让张量核、向量核与内存传输三类资源充分重叠工程层仓库的 sm90/mha.mojo、softmax.mojo、mha_utils.mojo 与 bench_flash_attention.yaml 完整落地了上述设计并在 design-docs/README.md 的目录中与其他设计文档如 matmul-to-flash-attention.md、uwgmma-flash-decoding.md共同构成 Modular Platform GPU 内核优化的知识体系。对于想继续深入研究的读者建议按以下顺序阅读先通读本文对应的 multi-head-flash-attention.md再结合 matmul-to-flash-attention.md 理解 FA 如何从快速矩阵乘法中演化而来最后对照sm90/mha.mojo的内核代码逐行印证 FA3 的 warp 特化与流水线结构。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考