一、砚溪镇有座理账楼,每日要对清六十四个商家的账

砚溪镇临河,镇上大小商家六十有四。商会底下有座理账楼,专管一件事:把六十四个商家两两之间的往来账,每日重新对清一遍。

六十四个对六十四个,便是四千零九十六格的偌大一张对账总表。哪两户之间收了多少、欠了多少、冲抵后净差几何,全在这张表里。

掌簿人姓简,人称简公。他手下新来了个学徒,叫桑丫,性子勤快,却还摸不透这楼的脾性。

头一日,桑丫按老法子干活:她把整张四千零九十六格的总表,从后仓的大案上整张扛到书房的小桌上,铺开,一格一格地算。

可那张表实在太大。后仓离书房百步,每次要核对一点什么,就得把整张表重新扛一遍。一日下来,桑丫肩膀肿了,账却只对了小半。

简公在廊下看了半晌,叹道:“你这法子,是把整座山都搬进屋来,再慢慢挑一粒沙子。”

二、简公说:桌小,便不要摊开整张表

桑丫委屈:“不摊开,怎知这一格该填什么?”

简公摇头:“账这事,要的是结果,不是把整张表都亮在眼前。你且看我的法子。”

他领桑丫到书房,指了指那张只能容八户账册的小桌。

“后仓的大案,能摊下整张表,可远、且慢,扛一次累死人。这书房的小桌,近、且快,可只放得下八户。咱们便认了这个短处——小桌小,就不摊整张表,只搬能放下的一摞来算。

“那四千多格,岂不永远算不完?”桑丫疑道。

“分摞搬。”简公笑,“把六十四户切成八摞,每摞八户。先搬第一摞的八户,与对面八户对账,八乘八,六十四格,小桌刚好容下。”

三、小桌上边算边记,不必回头看旧摞

桑丫照做,搬来头一摞,小桌上算出了这六十四格的局部小账。

“然后呢?”她问,“再搬下一摞,可这摞的结果,与上一摞怎么合?”

简公在纸上写下两个数:“你每算完一摞,只记下这两样——一是到目前为止见过的最大差额,二是所有已算格子加起来总共占的权重。下一摞搬来,你不必回头去翻上一摞,只拿这两个数与新摞的结果一并算,便把新结果并进了总账。”

桑丫试了:搬第二摞,与桌上那两个数一合,总账自动往前推了一步。再搬第三摞、第四摞……八摞搬完,她竟从未把整张表铺开过。

“那总表,”她忽然懂了,“从头到尾,从没真的存在过。”

简公颔首:“对。它只在你脑子里、在那两个数里,逐摞长成了最后的账本。书房干净,肩膀也不肿了。”

四、北风夜的一场乱账,法子经住了

入秋后有一夜,河上起了风,三家商号同时来改账,桑丫连搬了十几摞,其中还有重叠的户头。

若是从前,整张表摊着,改一处便要重扛整表,非乱不可。可这回,她只管一摞一摞搬,每搬一摞就更新那两个数、推进总账。重叠的户头,因为都归进同一套累计里,自然不冲突。

天亮时账清,分毫不差。

桑丫抚着小桌,念出一句话:“算得快的,不是桌子大,是懂得不把整座山搬进屋。”


技术解读

FlashAttention(Dao et al., 斯坦福, 2022;FlashAttention-2, 2023;FlashAttention-3, 2024)是一类IO 感知(IO-aware)的精确注意力算法,核心目标是绕开标准注意力在 GPU 显存带宽上的瓶颈。

标准自注意力要计算 N 个 token 两两之间的分数,得到一张 N×N 的注意力分数/权重矩阵。朴素实现会把这张矩阵**完整物化(materialize)**到 HBM(高带宽内存),再读回来做 softmax 与加权求和。当序列变长,N×N 矩阵呈平方级膨胀,既吃显存(O(N²)),又因反复读写 HBM 而受限于显存带宽——计算单元空转等待数据。

FlashAttention 的解法是 tiling(分块)+ online softmax(在线 softmax)+ 重计算(recomputation):把 Q/K/V 切成小块放进片上 SRAM,在 SRAM 内算局部注意力,用 running max 与 running sum 增量维护 softmax 归一化,只把最终输出 O 与统计量写回 HBM。整张 N×N 矩阵从不完整存在,显存降到 O(N),HBM 访问量大幅下降。

核心概念回顾

概念 通俗解释
注意力分数矩阵 N×N 每个 token 与所有 token 的关联强度,共 N² 项
HBM(高带宽内存) GPU 板上大容量但相对慢、离计算单元远的内存
SRAM / 片上内存 GPU 计算单元旁极小但极快的内存
物化(materialize) 把中间结果完整写进 HBM 暂存
tiling(分块) 把大矩阵切成能放进 SRAM 的小块分别计算
online softmax 边算边用 running max/sum 增量维护归一化,无需看全表
重计算(recomputation) 反向时重算中间量,而非从 HBM 读回,省显存
IO-aware 算法设计显式考虑“数据在快慢内存间搬动的代价”

故事中的隐喻对照

故事元素 映射的技术概念 解释
砚溪镇六十四户商家两两对账 N 个 token 的 N×N 注意力分数矩阵 每两户一对,恰如每个 query 与每个 key 配对打分
后仓大案(远、慢、能摊整表) HBM 高带宽内存 容量大但离计算远、读写慢,搬整张表代价高
书房小桌(近、快、只容八户) SRAM / 片上内存 极快但极小,放不下整张表,只能放一块
桑丫整张表扛来铺开 标准注意力物化完整 N×N 矩阵到 HBM 为算一格要把整表反复搬,慢且占显存
把六十四户切八摞、每摞八户 tiling 分块(block size) 每次只取能进 SRAM 的 Q/K/V 子块
小桌上算局部六十四格小账 在 SRAM 内计算 block 的局部 attention 局部打分与加权在快内存里完成
记下“最大差额”与“总权重”两个数 online softmax 的 running max m 与 running sum l 增量维护归一化所需的统计量,无需回看旧块
搬新摞不回头翻旧摞,并入总账 online softmax 增量更新,避免重算全表 新块结果用已有统计量并入,数学等价且不重算
八摞搬完、整张表从未铺开 不 materialize 完整 N×N 矩阵 矩阵只在累计中“虚拟存在”,显存 O(N)
只交回最后的总账本与两个数 只把输出 O 与 softmax 统计量写回 HBM 大幅减少 HBM 读写,绕开带宽瓶颈
北风夜重叠户头不乱 增量累计天然处理分块重叠,结果精确 FlashAttention 是精确注意力,非近似

为什么这个故事对应 FlashAttention?

  1. 因为大仓库搬整张表很慢、且占满场地,所以标准注意力把 N×N 矩阵物化到 HBM、反复读写既慢又吃显存——分小块在快桌上算的思路直接对应 tiling 提速。
  2. 因为书房小桌放不下整张表,所以SRAM 容量远小于 HBM,注意力必须切块计算(tiling),不能一次性摊开。
  3. 因为不铺开整张表就拿不到全局归一分母,所以需要 online softmax(running max + running sum)在每一块算完后增量维护归一化,而不必回看已算块。
  4. 因为只在书桌上算完交回总账与两个数,所以FlashAttention 只把输出 O 与统计量写回 HBM,显存从 O(N²) 降到 O(N)。
  5. 因为边算边并入是严格的数学等价变换(指数和的分解恒等式),所以FlashAttention 得到的是精确注意力,结果与传统实现一致,只是更快更省——不是近似算法。
  6. 因为分块计算把“反复搬大矩阵”换成“多次搬小块”,所以HBM 带宽瓶颈被绕开,GPU 计算单元利用率(MFU)显著提升,长序列下收益最大。

后记
算珠楼的道理说来简单——桌子小,就别把整座山搬进来。FlashAttention 做的正是同一件事:它没让 GPU 变得更大,只是教会计算单元“分摞搬、桌上算、边算边记”,于是那些曾经必须摊开在慢内存里的庞然大物,最终只在两个累计数里悄悄长成了答案。下次你看到模型能塞下更长的上下文,或许该想起砚溪镇那张从没真正铺开过的对账总表。