系统工程FlashAttention
从朴素实现到 Auto-Tuning
会员专享编写第一个 Flash Attention Kernel,并利用 Auto-Tune 进行性能优化。
在权益中心获取代码我们在上一章推导出了 Flash Attention 的数学公式(Tiling + Online Softmax)。现在,是时候把数学变成代码了。
核心循环结构
从公式到 Kernel,中间隔着三个必须先想清楚的问题。
第一,两层循环谁并行、谁串行。 上一章的推导写成了 "Outer Q, Inner K":外层遍历 的分块,内层遍历 的分块。但在 Triton 里你不会看到外层循环,因为它被 SPMD 模型隐式并行化了:每个 Program 认领一个 块,tl.program_id(0) 就是它的编号。真正写成 for 的只有内循环。这一点如果没建立起来,读 Kernel 时会一直找不到外层循环在哪。
第二,哪些数据在循环中不变。 块一旦载入就固定不动,整个内循环都在复用它; 则每轮重新加载。这个区分决定了 tl.load 写在循环外还是循环内,也是 Flash Attention 能把 HBM 访问从 降到 的直接原因。
第三,累加器要带几个状态。 朴素 softmax 可以先算完整行再归一化,分块之后不行:每读入一个新的 块,最大值可能被刷新,此前累加的结果需要按比例修正。所以除了输出累加器 acc,还要维护 running max 和 running sum 两个统计量,共三个跨迭代状态。
理清这三点,下面的 Kernel 就只是把它们逐条翻译成 Triton 语法:
登录以继续阅读
这是一篇付费内容,请登录您的账户以访问完整内容。
CookLLM文档