【CS336】lecture4 FlashAttention

FlashAttention概览

fa是个很经典的优化,和torch的标准注意力实现相比,访存量减少了一个数量级,运行时间也大幅降低,只有运算量稍微增加了一点,但正如我们前面提到的,GPU的运算吞吐远大于内存带宽,因此计算量增加问题不大,只要优化了搬运量就行,而这就是fa的核心思路。

回忆attention

回忆一下attention计算,QK先乘,再做一个逐行的softmax,得到注意力矩阵,再乘上V矩阵,得到注意力输出

tiling

那么和GEMM的优化类似,也可以先做一个QKV矩阵的tiling

朴素softmax

接下来是关键,GEMM已经被充分优化了,fa的创新点主要在于优化了softmax算子。softmax是个逐元素算子,计算强度很低,属于访存瓶颈。

先来看看朴素的softmax怎么实现,这里在实现上采用的是safe softmax,在原始定义的基础上,找到最大指数,所有指数先减去这个最大指数,再求幂。这样的好处是把所有的指数都映射到负数了,计算幂之后值域也在0-1之间,不会出现指数爆炸。

这并不影响答案的正确性,softmax本身是一个幂和所有幂之和的比值,这里对所有指数都减了mx,相当于对所有幂都除以一个常数emxe^{mx}emx,分子分母同除常数,结果不变。

为了实现这个思路,需要三次循环。

  • 第一次找到指数最大值mx
  • 第二次求所有指数减去mx后,以e为底的幂的和sum
  • 第三次利用sum求出每个位置的最终softmax输出

在线softmax

朴素softmax的访存太多了。fa的核心优化点是使用了在线softmax,把循环次数优化到两次

  • 第一次在线地求出到目前为止的最大值mx,以及这个最大值对应的sum,如果最大值变化了,就对旧的sum做一下缩放
  • 第二次循环利用sum和mx计算最终答案

对于softmax这个计算强度很低,几乎完全是memory bound的算子来说,减少一轮循环的访存可以带来巨大的提升。具体来分析,朴素3pass会3次读,1次写,在线softmax 2pass会2次读,1次写,访存量的比值是4/3=1.33,理论上的极限加速为1.33x

这看起来不多,但softmax在运行时间里占比不少,且之前被认为是不可优化的,现在把他也变成可优化了,开辟了一块以前没探索过的加速空间。

算子融合

当然能实现最开始图中的加速比,最核心的还是算子融合,减少数据的反复搬运,回到最开始这张图

HBM也就是显存的读写量,只有之前的10%,因为把多个算子拼成的attention算子改成完全融合了,这才是能得到5x-6x加速比的核心原因。

相关推荐
一条大祥脚15 小时前
【CS336】lecture4 MoE|专家路由|专家数量
moe·deepseek·cs336·专家路由
gravity_w20 天前
【CS336】Lecture 2 PyTorch 与资源核算(Resource Accounting)
人工智能·深度学习·语言模型·llm·nlp·flops·cs336
thesky1234561 个月前
27届大模型面试准备(五十九):大模型推理引擎内核深度剖析——PagedAttention、调度器与显存管理
大模型·vllm·flashattention·推理引擎·pagedattention·连续批处理·显存管理
thesky1234562 个月前
27届大模型面试准备(二十九):长文本推理与高效注意力——FlashAttention、稀疏/线性注意力与推理侧长度外推
大模型·面试准备·flashattention·推理优化·稀疏注意力·长文本推理·长度外推
nuowenyadelunwen4 个月前
CS336 Assignment 1 BPE分词器训练初版(朴素版基础上优化)及后续优化方向分析
llm·cs336
minhuan4 个月前
FlashAttention、PagedAttention两代注意力算法,改写大模型推理生态详解.186
自注意力机制·大模型应用·flashattention·pagedattention·注意力算法详解
Allenlzcoder5 个月前
Stanford CS336(2026)课程介绍
cs336
墨心@6 个月前
Byte-Pair Encoding (BPE) Tokenizer
人工智能·自然语言处理·nlp·datawhale·cs336·组队学习
爱听歌的周童鞋7 个月前
斯坦福大学 | CS336 | 从零开始构建语言模型 | Spring 2025 | 笔记 | Course Summary
llm·cs336·course summary