2026/10/10 4:30:42

FlashAttention在AMD MI50上的工程落地:数据流优化、在线softmax与验证实践

FlashAttention在AMD MI50上的工程落地:数据流优化、在线softmax与验证实践 这个题目我琢磨了几天才动手写。之前几篇笔记要么在讲算子层面的编译优化要么在讲显存管理这篇轮到FlashAttention了。但我不想再写一篇什么是FlashAttention的科普文网上已经有太多论文笔记了。我更想聊的是当你真的要把FlashAttention从纸面公式变成能在老一代GPU比如AMD MI50上稳定跑出性能的kernel时那些数据流细节和工程验证方法才是真正卡脖子的地方。先说下这篇文章的适用人群如果你正在做AIInfra相关的工作或者要在大模型训练/推理场景里手工优化attention算子又或者你需要在ROCm生态里做算子移植那这篇笔记应该能给你一些参考。我不会贴大段论文公式重点是讲清楚数据流为什么这样走、在硬件上如何落地、以及怎么验证它真的没算错、真的变快了。1. 显存墙与计算墙标准Attention的数据流困境1.1 一次矩阵乘法就把显存打爆的N×N矩阵先说一个我在实际项目中反复遇到的场景。当序列长度N从512涨到4096、再到8192很多原本能跑的模型突然就OOM了。如果你去排查会发现罪魁祸首通常不是模型参数而是attention中间那个S矩阵。我们拿一个典型的配置来算算batch1、heads32、seq_len4096、head_dim64。标准attention计算S Q K^T每个头的S矩阵是4096×4096的浮点数占64MB用FP32的话32个head就是2GB。这还只是单层transformer。一个12层的模型光attention的中间矩阵就要吃掉24GB显存。这还没算softmax之后产生的P矩阵以及反向传播要用的梯度矩阵。在一个16GB的MI50上单卡想跑个12层模型加4096序列长度直接就没戏了。问题的本质在于S矩阵的规模是O(N²)的而显存是有限的。更麻烦的是这个矩阵不仅在显存里占地方还会被反复读写。标准attention在GPU上通常分成三步先算S QK^T把结果写回HBM这里泛指显存再做softmax读S、写P最后算O PV读P。每一步都是一次完整的显存读写往返。还是上面那个例子S写一次读一次就是4GB流量P再写一次读一次又是4GB加上Q、K、V、O本身的搬运单层attention一次前向就要触碰十几GB的显存流量。这里有个反直觉的点attention的FLOPs看着不多但慢就慢在数据搬运上。按N4096、head_dim64算单层attention大约2×N²×d×heads ≈ 2×16.7M×64×32 ≈ 68GFLOPs听起来不小但在能跑十几TFLOPs的GPU上理论上几毫秒就能算完。实际却慢一个数量级因为算力大部分时间在等数据从显存搬过来。所以从工程角度说标准attention是典型的memory-bound算子你优化它本质上是优化显存访问而不是省计算。1.2 Kernel融合的目标把三次访存压成一次理解了上面的瓶颈优化思路就清楚了不要让中间矩阵落地。理想情况是把Q、K、V读进片上SRAM之后在寄存器/共享内存里完成S的局部计算、softmax、以及乘V最后只把O写回显存。但这里有个数学上的拦路虎softmax需要全局的max和sum。标准做法是先扫一遍S找全局max做减max的指数归一化得到P再乘V。这个两遍扫描天然要求你知道完整的S无法边读边算。于是就有了FlashAttention的核心思想——在线softmaxonline softmax用分块的方式边走边维护running max和running sum把两遍扫描压缩成一遍。这样GPU kernel才能做到一次融合把S矩阵彻底消灭在片上。从数据流角度总结一下FlashAttention做了什么它把attention的计算从N×N矩阵的显存往返变成了N×d向量块的持续流动。每个tile的K、V从显存只被加载一次理想情况下S和P这些中间矩阵要么分散在SRAM里要么根本不存在于显存中。这一改显存占用从O(N²)降到O(N)访存流量也大幅下降。你去看FlashAttention论文里的图三角形和梯形的分块示意本质上就是在讲这个数据流变化。2. Online Softmax分块计算里的数值核心2.1 从两遍扫描到一遍扫描的在线合并FlashAttention能成立数学地基是online softmax。我先用一个不严谨但好懂的方式说明白它干了什么。普通softmax在只有一个block时很好算但attention必须分块。假设Q被切成两块Q₁和Q₂K、V也切成对应的块。你先用Q₁和所有K块算出一部分S做softmax得到输出O₁这个过程你维护了第一次遇到的max和sum。接着处理Q₂时对同一个query行你可能发现了更大的max值。这时候之前已经算好的那些exp(s - max)如果还按旧max归一化就过期了需要乘以一个缩放因子exp(old_max - new_max)来修正。关键是这个修正可以沿路累积——每次发现新max就把旧sum乘个系数调整。换句话说你不需要见到完整的S行也能精确算出最终的softmax结果。每个分块只带着自己的running max和running sum向前流动这就是online softmax的工程本质。2.2 前向循环里的块调度与统计量更新放到实际kernel里看循环结构。外层循环遍历Q的每一个分块比如128行对于每个Q分块内层循环从头到尾遍历K、V的分块。这个内层循环里发生的操作序列是固定的乘出S分块一个小的局部矩阵、对比当前分块的局部max和之前的running max、更新running sum、计算局部P、累加输出O。我建议用一个小例子感受这个合并过程。假设某一行S实际是[1, 2, 3, 4]但分块来了先见到[1, 2]max2sumexp(0)exp(1)≈1.368不对这里为了简化我直接说思路以2为基准局部sum exp(1-2)exp(2-2)0.36811.368然后乘上V的对应块得到部分O。再见到[3, 4]新max4比2大。这时候之前那些exp值的基准从2变成了4每个都要整体缩小exp(2-4)exp(-2)≈0.135倍。所以running sum要乘以0.135加上新的exp部分exp(3-4)exp(4-4)0.36811.368总sum变成1.368×0.1351.368≈1.553。输出O也要整体乘以同样的缩放因子。这样一路合并直到跑完所有K、V块得到的O和精确softmax结果一致。这个过程中有个容易忽略的点缩放因子的累积几乎不带来额外数值误差。因为每一步你都在做数学上等价的操作只是顺序不同。真正决定精度的是你用什么精度存S分块和P分块这涉及到下一个话题。2.3 数值稳定与Float16累加误差的平衡很多人在工程实践里纠结FlashAttention用FP16到底行不行我的经验是核心运算用FP16会损失一些精度但通过合理的分块大小和数据尺度管理误差可以控制在训练可接受范围内。FlashAttention不需要像普通softmax那样减全局max来避免溢出。因为它用的是running max机制——每个分块内先减掉当前的局部max而这个局部max是逐步逼近全局max的。所以exp的指数始终保持在合理范围不会溢出除非这一行本身就存在NaN/Inf。这一点很多教程没讲清楚容易让人觉得online softmax只是省了一次扫描其实它还顺带解决了数值稳定性。真正需要注意的误差来源是FP16的exp累加。FP16能表示的范围有限精度大约只有10位左右。在长序列下running sum累加的值可能非常大小的加数会被吃掉。实践中我通常的做法是S分块和P分块用FP16计算但running max、running sum、以及O的累加用FP32完成。代价是多占一点共享内存和寄存器但换来的数值余量很大。还有个细节SRAM里一个FP32的tile占的空间是FP16的两倍这会直接影响你能选的块大小。如果你的kernel最终在数值验证阶段误差超标多半不是online softmax的算法问题而是FP16累加的位置和精度没分配好。这时候优先把sum提升到FP32而不是一上来就切换整个kernel到FP32。我们后面第五章讲验证时会再展开如何量化这个误差。3. 反向重计算省显存的另一半秘密3.1 不存S矩阵梯度怎么算FlashAttention省显存不只是前向的事反向才是重头戏。标准attention反向需要用到P矩阵和softmax的中间梯度因此前向时你必须把S或P留在显存里。但FlashAttention反向上的做法很激进前向只保留每块的小规模统计量running max和running sum尺寸是N×d级别的以及输出O、QKV本身。反向需要P时不读显存而是重新按同样的分块方式算一遍前向在线重算出P分块。这样做的直接收益是反向不需要额外存储任何N×N的中间矩阵显存占用不随序列长度平方增长。代价是反向多了一次前向的计算量等于是用FLOPs换显存。在当前的硬件条件下计算多出来的部分通常远小于访存节省的部分所以总体提速反而是正的。3.2 梯度的计算公式与重计算的数据流反向传播的数据流看起来是这个样子的。你先从损失拿到dO为了算dV你需要P^T dO为了算dK和dQ你需要dS。dS的公式里包含两项一个是P逐元素乘以(dO V^T)另一个是减去一个行向量softmax的梯度修正项这个修正项等于P每行与(dO V^T)每行的点积然后按行广播。这一串公式里到处都需要P。于是反向kernel的内层循环退了回去外层遍历Q块内层遍历K、V块在块级别上重新计算S块、根据缓存的统计量还原P块然后立刻参与dV、dK、dQ的累加。P块在寄存器里算出来、用完即弃完全不落地显存。从数据流优化角度看这个重计算模式的价值甚至比前向还大。反向如果存P不仅要占O(N²)显存还要额外读一遍P而重计算虽然让S块的乘法多算了一次但P块完全不用访问HBM整体的访存流量反而更少。访存减少就是memory-bound算子提速的根本。3.3 重计算的代价与收益到底怎么算有些同学看到反向要多算一次前向就担心性能。我们来算一笔账。假设反向分两次循环第一遍先跑一个轻量的循环算出每行的softmax梯度修正项这需要dO和V第二遍再走完整的分块重计算循环把三个梯度都累加完。这样反向总的矩阵乘法次数大约是一次dOV^T、一次P^TdO、一次dS^TQ或类似、加上重计算S块的那次QK^T。你比标准反向多了一次QK^T的FLOPs但省去了O(N²)矩阵的读写。对于batch8、heads32、seq4096这种规模省下的显存是几十GB量级而多出来的FLOPs是几十GFLOPS量级。aos跑在1TFLOPs以上的GPU上这点额外计算可能只占百分之几的时间但省下的几十GB显存意味着你能把batch翻倍、或者把序列长度翻倍这是质的差别。所以重计算这个设计在工程上几乎是纯赚。4. Composable Kernel与MI50适配没有Tensor Core怎么办4.1 MI50的硬件边界决定优化方向新闻热搜里提到的ck flashattention 适配mi50中的ck应该是指AMD的Composable KernelCK库。要在这块卡上移植FlashAttention首先得搞清楚MI50的硬件边界。MI50是Vega 20架构GCN时代的产物。这意味着几个关键特性。其一它没有CDNA架构上那些针对矩阵运算的MFMA指令matrix core这类东西在MI50上是不存在的。你做矩阵乘法得靠普通VALU指令或者GCN上还算好用的vector dot积指令FP16的v_dot2可以一次算两组半精度乘加。因此矩阵乘这部分没法用MFMA快速实现tile的尺寸选择会受到寄存器压力的严格限制。其二MI50的wave是64 laneswave64不像NVIDIA是32线程的warp也不像CDNA的wave32模式。这意味着你的访存指令一次能覆盖更多数据分支发散的影响范围也更大。其三每CU的LDS共享内存是64KB寄存器文件是每CU 256KB。这决定了块调度策略LDS里能放两个128×64的FP16块Q块和K块各一个V块循环覆盖再加上P结果块和O块基本上是紧巴巴的设计时必须精确算好LDS预算。4.2 为什么选Composable Kernel而不是手写HIP在做FlashAttention移植时我们有两条路一条是直接用HIP写一个融合kernel另一条是用Composable Kernel。我个人推荐CK尤其当你需要在多个GPU架构上维护算子时。CK的核心思想是把GPU kernel拆成非常细小的算子原语比如矩阵乘法的分块、布局转换、LDS读写、寄存器级的数据搬运然后通过模板元编程在编译期把它们拼装成完整的kernel。这让kernel的实现变成了一种数据流描述你告诉CK数据从哪里来global memory、中间放哪LDS mutate、怎么算MFMA或普通乘加、结果怎么走回归global。CK负责把这些组合成高效的汇编级代码。对FlashAttention来说CK的价值是提供了一个ThreadblockSwitcher之类的调度框架以及现成的tile遍历结构。你不需要自己写循环展开、不需要手动处理LDS的padding和bank conflict虽然还是需要理解能把精力集中在attention特有的online softmax逻辑上。而且CK的GPU抽象天然考虑了MI50这种不带MFMA的硬件它会自动生成fallback的VALU路径。如果纯手写HIP大概率会踩进寄存器分配和指令调度的坑里调起来非常痛苦。4.3 适配MI50时真正影响性能的几个点这块我直接给结论都是实践中踩出来的。第一tile size不要照搬论文或CUDA实现。FlashAttention论文中常用的是64×64或128×128但那是针对NVIDIA的Tensor Core设计的。MI50没有Tensor Core寄存器做矩阵乘时过大的tile会导致寄存器溢出反而变慢。我调下来Q块取128×64、KV块取64×64的方案比较合适。这样LDS里同时放一个Q块128×64×2B16KB和一个K/V块64×64×2B8KB再加上P/Pacc的缓冲LDS用量大概在40KB左右64KB的预算里还留了余量。第二注意128-bit向量访存的对齐问题。MI50的global load吞吐最高时要求128-bit16字节对齐的连续访存。如果数据布局是标准的(B, H, N, D)N×D的最后一维是64个FP16正好是128字节天然满足128-bit对齐。但如果你的QKV已经在之前的算子中被重排过、有offset偏移就很容易出现misaligned access性能暴跌一半以上。我在调试时专门写过一个对齐检查的单元测试把所有可能进入attention的tensor布局都跑一遍避免上线后才发现隐藏的对齐问题。第三处理LDS的bank conflict。Vega的LDS有32个bank每个bank 4字节。当你按(64, 64)的FP16 tile去读取时同一行64个元素分布在连续的128字节中正好覆盖32个bank的两轮如果没有padding第二个半行会跟第一个半行冲突。解决办法很简单在分配LDS时给每行加上4字节的padding让行首偏移不再是32的倍数。这个改动在实测中能带来10-15%的访存性能提升。第四wave64的分支发散影响必须提前考虑。如果是64 lanes一个wave而你的tile规模刚好是128那么一个wave内的两个半部分如果走了不同的分支比如padding mask整个wave的执行效率会很低。所以对变长序列的mask处理不要在每个元素级别做判断而是把mask预先压成block-level的标记在块循环开头做一次分支块内部走无分支路径。这个小优化在动态batch推理场景下收益特别明显。5. 工程验证链路从数值正确性到有效带宽5.1 基准对照怎么证明你没算错一个高性能kernel如果算错数跑得再快也是废的。所以工程验证第一步永远是搭一个可信的golden reference。我用的是很朴素的办法写一个PyTorch的eager attention版本用FP32计算把每一层的S、P、O都导出来作为基准。这样做的原因是FP32的数值在常规序列长度下几乎不会有大误差拿它做标准可信度高。弘一的误差指标我建议用两个一个是最大绝对误差max abs error直面极端错误另一个是均方根误差RMS error反映整体偏离。我实测下来FP16的FlashAttention与FP32的eager版本对比最大绝对误差通常在1e-2到1e-3量级RMS误差在1e-3到1e-4量级。训练场景下这个精度是可接受的因为后续的梯度下降本身对噪声有容忍度。但如果你的误差到了1e-1级别那就有问题了优先检查是否online softmax的缩放因子没及时更新。测试用例不能只用标准正态分布那是理想情况。我通常会加这几类全相同值比如所有元素都是1这能测出softmax归一化是否正确极大值与极小值混合比如1e4和-1e4测exp溢出路径以及随机mask模式测padding部分是否被正确屏蔽。这四个用例跑下来数值逻辑有没有大问题基本能现形。5.2 性能测量有效带宽才是硬指标性能验证环节我很少只看kernel耗时而是计算一个更能说明问题的指标有效带宽effective bandwidth。对FlashAttention这种memory-bound算子有效带宽表示kernel从显存读取的数据量与耗时的比值。公式可以粗略写为有效带宽 (读取的Q K V 写出的O字节数) / kernel耗时比如N4096、head_dim64、heads32、batch1时QKV和O的字节数加起来大概是(4096×64×32×3×2B 4096×64×32×2B) ≈ 80MB。如果耗时是0.2ms有效带宽就是400GB/s。MI50的HBM2理论带宽约1TB/s达到50-60%就算合格。我调好的kernel在合理配置下能达到550-650GB/s左右的有效带宽不同驱动的版本会有浮动距离理论极限还有空间但已经接近访存密集算子的正常水平。计时时要注意几个坑一是必须先warmup几轮让GPU频率和缓存状态稳定二是不要在hipEvent计时区间里混入host到device的内存拷贝那是另一笔开销三是跑多个batch取平均单次调用抖动太大。我自己的标准流程是warmup 10次、正式测50次、取p95和median两个值。只看mean容易被极端值带偏。5.3 显存水位与长序列稳定性除了单纯的正确性和速度工程上还有两个必须验的项目显存峰值和长序列稳定性。显存验证其实简单粗暴写一个脚本把序列长度从1024依次加到8192或更高分别在eager attention和FlashAttention下训练一个小模型记录峰值显存。你会看到eager曲线的斜率是二次方的FlashAttention则是近线性。比如Seq4096、12层、batch4时eager版本已经逼近20GB而FlashAttention可能只用了6GB。这种对比发出来非常直观也是你向团队证明投入产出比的最好依据。长序列稳定性指的是随着N增大online softmax的累积误差会不会漂移。我测试的结论是只要running sum保持FP32FP16的FlashAttention在N加大到8192、甚至16384时RMS误差也基本稳定在同一量级没有明显劣化。但这个结果依赖硬件和编译器版本换一张卡、换个ROCm版本就得重新测一遍。这也是为什么工程验证不能只做一次要固化成一个回归脚本每次改kernel或换环境都要跑一遍。6. 调试过程中踩过的坑与实际建议6.1 最隐蔽的bug分块边界与mask处理调试FlashAttention这类kernel最常见的隐性bug不在算法公式而在分块边界。比如序列长度N不是tile size的整数倍时最后一块K、V的处理如果你直接把越界部分读进来就会读到下一段内存虽然多数时候不会段错误但会引入垃圾值导致那一行的输出完全错乱。解决办法是在内层循环末尾对最后一组tile做边界裁剪tile上的元素级mask只计算真实存在的部分。这个逻辑在纯Python端很好写但到了kernel的汇编级每个lane要对齐处理非常容易写错。我建议把它提取为独立的函数并针对几种边界条件专门写单元测试。另一个隐蔽坑是因果maskcausal mask。很多模型attention是带因果性的Q的第i行只能看到K的前i行。如果直接把mask值设为0而不是-Infsoftmax的归一化分母里就会混入无效位置最终输出偏小。我见过不止一个人的第一次实现在这里算错表面看pytorch对比误差不大但训练曲线会异常。正确做法是把mask位置的分数直接设为负无穷或单独维护一个有效累加计数器。6.2 性能回退时的定位思路当你发现kernel比预期慢不要盲目去调tile size。我一般的排查顺序是先确认是否命中了正确性回归如果数值都是错的性能数据没有意义然后用ROCm的profiler比如rocprof看kernel内部的时间分布到底是global load阶段占大头、还是计算阶段占大头再检查有没有意外的本地内存溢出spill。有一次性能回退的案例特别典型我把tile从64×64换成128×64后理论访存量更低但实际反而慢了15%。用rocprof一查发现寄存器溢出导致本地内存访问量暴涨。原因是128×64的Q块在寄存器里的转置拷贝比64×64多了一倍寄存器不够用打到了local memory。解决很简单把Q块在LDS里提前转置好而不是在寄存器里转。这个经验说明一个道理在GCN这类没有Tensor Core的硬件上寄存器的调度比计算本身更影响性能任何看着计算变少的优化都可能被寄存器溢出抵消。6.3 如果让我重新做一遍的顺序回到开头那个热搜词适配MI50如果要做类似工作我建议的顺序是先花一天时间把online softmax的数学和分块循环在图层面写通可以先用Python numpy模拟不急着写kernel再花一天写一个最小的HIP kernel在单层上验证正确性然后才用CK重构到生产级。正确性路径走通之前不要碰任何性能调优参数。性能调优阶段每改一个tile size或bank padding就把正确性回归测试再跑一遍防止为提速而牺牲正确性。我自己吃过跳过小模型模拟的亏直接在kernel里写online softmax合并逻辑结果遇到了一个特别难查的sum缩放问题。后来回到numpy模拟五分钟就找出了逻辑漏洞。所以别嫌先模拟这一步废时间比起在汇编级别debug数值逻辑这点前期成本低太多了。另一个建议是如果你有CUDA FlashAttention的参考实现不要直接照着搬逻辑而是要理解它为什么那样写。比如CUDA实现里常用warp shuffle来做跨lanes的max归约但MI50是wave64CUDA的warp是32shuffle的语义和范围都不一样。照搬代码或照搬思路都会踩坑。一定先从数学上确认哪里需要归约沟通再根据wave64的特性重新设计归约方式比如用shared memory做两级归约。最后这份适配工作做完之后把验证脚本和perf基准固化成一个可重复的CI任务。我见过太多团队优化完算子就撒手不管了结果下次驱动升级或换一批卡性能莫名其妙变差却没人知道为什么。固化下来的回归系统后续带来的长期回报往往比最初那两三个月的优化本身还要大。