LEARNING UNIT · 12
在线 Softmax 与 FlashAttention
用分块等价性和在线归约状态解释 FlashAttention 的前向、反向、掩码与显存复杂度。
- 已整理章节
- 12 节
- 单元来源
- 11 条视频
- 总时长
- 32:04
- 状态
- 已发布
- 学习位置
- 12 / 20
主题讲解 · 02:32
为什么分块计算仍得到同一个注意力矩阵
学习目标
- 能从单个元素 解释 的含义。
- 能写出 Q 与 K 分块后每个输出块的公式。
- 能证明分块只改变计算顺序,不改变注意力分数矩阵。
- 能说明 HBM、SRAM 与线程块在分块计算中的角色。
- 能指出“矩阵乘法等价”与“完整 FlashAttention 正确性”之间还差什么。
前置与衔接
需要会做矩阵乘法,并知道朴素 attention 的分数矩阵来自 。
本课是“在线 Softmax 与 FlashAttention”单元的入口。先证明 tile 计算不会改写 ,再学习在线 softmax 如何在不保存完整分数矩阵时得到同一输出。
核心讲解
1. 朴素算法先物化完整分数矩阵
忽略缩放因子 ,注意力分数为
朴素注意力先计算完整的 A=QKᵀ,示例矩阵给出可逐项核对的结果。
原视频 · 00:20 ↗若 ,则 。
朴素实现通常把这个完整 中间量写到高带宽显存 HBM,再读取它做 softmax 和后续乘 。
当序列长度 增大时,这个中间量的元素数按 增长。
FlashAttention 的核心工程目标不是修改 attention 定义,而是减少这个中间量在 HBM 中的读写和物化。
2. 每个输出元素本来就是一个局部内积
矩阵乘法定义给出
注意力分数 Aᵢⱼ 等于 Q 的第 i 行与 K 的第 j 行做内积。
原视频 · 00:40 ↗这里 是 Q 的第 行, 是 K 的第 行转置成的列向量。
展开后
这条公式非常重要:任何分块方案只要让每个 最终累加完全相同的 个乘积,结果就不会改变。
所以正确性检查不需要迷信“FlashAttention”这个名字,只要回到每个输出元素的求和项即可。
3. Q 按行切,Kᵀ 按列切
把 Q 按 token 行切成两个块:
K 原本也按 token 行切:
转置后就成为列块:
Q 沿行方向切块,Kᵀ 沿列方向切块,二者组合覆盖完整的注意力矩阵。
原视频 · 01:00 ↗因此
每个 都是完整注意力矩阵中的一个矩形 tile。
4. 拼回所有 tile 就是原矩阵
课程用具体数字逐块计算,发现每个 tile 与朴素结果中相同位置的小块完全一致。
每个输出块 QᵢKⱼᵀ 仍由相同的行列内积组成,拼接后与朴素结果一致。
原视频 · 01:20 ↗原因并不神秘。
设第 个 Q 块覆盖行集合 ,第 个 K 块覆盖行集合 。
块乘法 中任意元素仍是
这与朴素矩阵中 的定义一字不差。
分块改变的是:
- 先算哪些行列区域;
- 一个 tile 何时搬进片上存储;
- 中间结果是否立即写回 HBM。
分块没有改变的是:
- 每个输出元素对应的 Q 行与 K 行;
- 内积的收缩维度 ;
- 每个内积包含的乘加项。
因此,在相同浮点精度与舍入顺序假设下,分块结果与朴素结果相同;实际浮点实现可能有极小舍入差异,但算法语义等价。
5. GPU 上为什么值得分块
HBM 容量大但访问代价高,片上 SRAM 容量小但带宽高。
FlashAttention 让线程块负责一个或若干 tile:
- 把 Q tile 从 HBM 搬到 SRAM;
- 分批把 K/V tile 搬到 SRAM;
- 计算局部 ;
- 立即把局部结果并入 softmax 状态和输出累加;
- 不把完整 分数矩阵长期写回 HBM。
GPU 线程块把 Q/K tile 从 HBM 搬到 SRAM 计算局部注意力块,而 softmax 仍需要跨块的全局行信息。
原视频 · 02:20 ↗课程在结尾指出一个关键困难:softmax 是逐行的,但每个线程块一次只看到一行的一部分。
这意味着证明 分块等价还不够,还要证明局部 softmax 状态可以正确合并。
6. 编者补充:完整正确性还需要在线 softmax
对一行分数 ,稳定 softmax 通常使用全局最大值
和归一化分母
若当前只看到一个分块 ,可先得到局部最大值 与局部分母 。
把旧状态与新块合并时,先更新
再把二者缩放到同一基准:
输出的加权和也做同样重标定,最终再除以 。
因此完整 FlashAttention 正确性包含两层:
- 分块矩阵乘法覆盖同一组注意力分数;
- 在线 softmax 用可合并状态恢复与全局 softmax 相同的归一化输出。
本条视频完整说明了第一层,并在结尾引出第二层;它不是对整个 FlashAttention 前向算法的全部证明。
跟练与练习
原视频练习
编者练习
Q 有 8 行、K 有 12 行,二者特征维度都是 64。若 Q 每 4 行一块,K 每 3 行一块,最终会得到多少个注意力 tile?每个 tile 的 shape 是什么?
查看参考答案
Q 有 个行块,K 有 个行块,因此输出有 个 tile。每个 tile 是 ,shape 为 ;拼接后得到 的完整分数矩阵。
编者练习 2
为什么“每个输出 tile 都正确”仍不足以证明 FlashAttention 的最终输出正确?
查看参考答案
attention 还要对整行分数做 softmax,再乘 V。单个 tile 只包含一行的局部分数,缺少全局最大值和全局归一化分母。还必须证明在线 softmax 的最大值、分母和加权输出状态可以跨 tile 正确合并。
常见误区
- 误区:K 横着切,所以 Kᵀ 仍横着切。纠正:转置后 token 行块会变成列块。
- 误区:分块近似了 。纠正:它重排精确乘加,算法语义并非低秩或稀疏近似。
- 误区:FlashAttention 把注意力复杂度从 变成线性。纠正:计算量通常仍是二次级,主要降低的是中间显存与 HBM I/O。
- 误区:证明块矩阵乘法就完成了全部证明。纠正:还需要在线 softmax 与乘 V 的合并正确性。
- 误区:浮点结果必须逐 bit 相同。纠正:重排归约顺序可能产生微小舍入差,但应在数值容差内等价。
本课小结
- 始终是 Q 的一行与 K 的一行做内积。
- Q 行块与 Kᵀ 列块的笛卡尔组合覆盖完整注意力矩阵。
- tile 计算改变执行顺序与内存流量,不改变每个输出元素的定义。
- GPU 分块的收益来自复用片上 SRAM、减少完整分数矩阵的 HBM 读写。
- 下一步要学习在线 softmax 如何合并局部最大值、归一化分母与加权输出。
主题讲解 · 02:31
按 Q 行分块为什么仍等于完整注意力分数
学习目标
- 能从元素定义解释注意力分数矩阵。
- 能证明只沿 Q 的行轴分块不改变 。
- 能区分“结果完全相同”与“执行代价一定更低”。
- 能说明完整 key 轴为何提供一行 Softmax 的全局统计量。
- 能识别视频中 CPU/缓存示意与具体实现之间的边界。
前置与衔接
注意力在忽略缩放、掩码和多头维度后,可先写成
视频用一个称为 SlimAttention 的方案说明:把 沿 token 行切成 tile,但在当前计算层次让 保持覆盖完整 key 轴。
问题是,分开计算每个 query tile,最后还能否严格得到原来的 ?
本课先证明线性代数等价,再讨论这一分块对 Softmax 与内存层次意味着什么。
核心讲解
1. 元素级定义不随分块改变
设
则
第 个元素为
注意力分数矩阵的元素 等于第 个 query 行向量与第 个 key 行向量的内积。
原视频 · 00:20 ↗图中选择一行 和 的一列,也就是原 的一行 。
二者内积填回 。
只要分块后仍计算同一对 ,该元素的值就不会改变。
2. 沿 Q 的行轴切成两个 tile
把 按行分成
其中 、 分别包含若干完整 query 行。
视频中的 SlimAttention 方案沿 token 行把 切成多个 tile,而 在该层解释中保持整块。
原视频 · 00:40 ↗这里没有沿特征维 切断单个 token 向量。
每个 query tile 仍包含完整的 维向量,因此可以与每个 key 做完整内积。
3. 分块乘法的直接证明
矩阵乘法对纵向拼接满足
左侧是一次计算完整 。
右侧是分别计算两个行块,再按原顺序纵向拼接。
与 分别计算后按行拼接,恰好恢复完整的 。
原视频 · 01:20 ↗两者逐元素相同,不依赖近似,也不依赖浮点外的特殊假设。
在真实浮点执行中,不同 kernel 的归约顺序可能造成末位舍入差异;这里的“等价”指相同数学表达式,不承诺 bitwise identical。
4. 为什么不能把结果横向拼错
、 是沿行轴切分,所以各自产物 shape 为
它们都覆盖完整 key 列,只覆盖不同 query 行。
因此必须沿行轴拼接。
若错误地沿列轴拼接,shape 和 的 query/key 语义都会改变。
判断拼接方向时,应始终追踪“被切的是哪一根轴”。
5. 一个具体 shape 例子
假设
将 切成两个 tile:
则
按行拼接得到 ,与完整乘法 shape 一致。
视频板书中的数字例子还逐项验证了对应元素相同。
6. 从数学分块到执行分块
视频随后给出内存示意:Q、K 位于较慢的主存层,执行单元把当前 Q tile 与所需 K 数据搬入更快的片上存储,再计算一个 A 的行块。
一个执行单元把当前 query tile 与所需 key 数据从较慢内存搬到片上快速存储,再计算对应分数行块。
原视频 · 01:40 ↗一个线程或线程块负责哪个 tile、K 是否一次完整驻留,都取决于硬件与实现。
数学证明只要求每个所需内积最终被计算,不能从证明本身推出某种固定调度一定最快。
7. 为什么完整 key 轴对 Softmax 有利
注意力 Softmax 沿 key 轴逐行计算:
若一个 Q tile 与完整 相乘,所得行块覆盖每一行的全部 个 key 分数。
当前 query tile 覆盖完整 key 轴时,得到的是完整分数行块,可直接获得该行 Softmax 的最大值与分母。
原视频 · 02:20 ↗于是该 tile 内可以直接得到每行的全局最大值 与完整分母。
不必跨多个 K tile 逐次修正同一行的 Softmax 状态。
8. 与 FlashAttention 双轴分块的差别
FlashAttention 通常同时对 query 轴与 key 轴分块。
单个分数 tile 只看到部分 key,因此没有一整行的全局最大值与分母。
它需要维护在线 Softmax 状态,在新 K/V tile 到来时重标定旧的部分结果。
视频中的 SlimAttention 行分块则让一个 Q tile 在该层次覆盖完整 K 轴,所以行归一化更直接。
但完整 K 是否能放进快速存储、是否需要在更低层继续切块,是实现层问题。
9. 正确性不等于任何尺寸都高效
沿 Q 行分块永远保持矩阵乘法的数学正确性。
性能却取决于:
- 的 K 数据是否适合目标缓存层;
- 每个 Q tile 的大小是否产生足够并行度;
- K 是否被多个执行单元重复搬运;
- Softmax、V 乘法和输出写回如何融合;
- CPU/GPU 的缓存容量、带宽和向量指令。
因此不能从“线性代数保证正确”直接跳到“所有长序列都更快”。
10. CPU、DRAM、SRAM 的术语边界
视频画面用 DRAM/SRAM 表示慢速与快速存储,并口述 CPU cache。
在 CPU 上更常见的具体层级是 DRAM 与 L1/L2/L3 cache;在 GPU 讨论中常见 HBM、shared memory、register。
它们共同表达“分层存储与块复用”这个概念,但容量、编程模型和搬运机制不同。
课程稿保留视频的抽象,不把板书中的 SRAM 标签当成所有 CPU 实现的字面硬件接口。
11. 方案与版本边界
“SlimAttention”并不是一个仅凭名字就能唯一确定所有分块、缓存与 Softmax 细节的通用规范。
本课只解释视频所示的“Q 行分块、当前层次完整 K”路径。
阅读具体论文或代码时,还要核对:
- 是否真的完整保留 K,还是在下一级缓存继续分块;
- 是否处理 causal mask;
- 是否融合 Softmax 与 ;
- 面向 CPU 还是 GPU;
- 支持哪些 dtype 与序列长度。
跟练与练习
原视频练习
编者练习
设 、,把 Q 沿行切成大小 2、4 的两个 tile。两个局部乘积的 shape 分别是什么,怎样恢复完整 A?
查看参考答案
。两个 Q tile 的 shape 分别为 与 ,乘积分别为 与 。沿行轴纵向拼接后得到 ,即完整 。
编者练习 2
为什么沿 key 轴把同一分数行切成两块后,不能各自独立做 Softmax 再直接拼接?
查看参考答案
Softmax 的分母和最大值必须覆盖该 query 行的全部 key。两块各自归一化会让每块权重和都等于 1,拼接后不再是完整行的 Softmax。必须先合并全局统计量,或使用在线 Softmax 的重标定公式。
常见误区
- 误区:分块会近似矩阵乘法。纠正:按行分块并完整计算所有内积时,数学结果严格相同。
- 误区:Q 按行切后应横向拼接结果。纠正:切的是 query 行轴,结果也沿行轴纵向拼接。
- 误区:一个 tile 可以切断 token 的特征维却仍直接套同一证明。纠正:沿收缩维切分时还需要对部分内积求和。
- 误区:完整 K 意味着所有实现都把整个 K 一次放进 SRAM。纠正:视频描述的是当前抽象层,实际可能有下一级分块。
- 误区:有完整分数行就完全没有缓存成本。纠正:仍需考虑 K 搬运、Q tile、输出和硬件容量。
- 误区:数学等价必然带来性能提升。纠正:性能还取决于数据复用、带宽与并行度。
本课小结
- 给出分块正确性的元素级依据。
- 将 按行切分,有 。
- 当前 Q tile 覆盖完整 key 轴时,可直接获得每行 Softmax 的全局统计量。
- 分块把数学计算映射到分层存储,但正确性与性能是两个结论。
- 视频的 SlimAttention、CPU 与 SRAM 表述属于具体示意,实际实现需按论文、代码和硬件复核。
主题讲解 · 03:44
在线 Softmax 如何把安全计算从三遍降到两遍
学习目标
- 能区分朴素 Softmax、安全 Softmax 与在线 Softmax。
- 能解释减去最大值为何不改变概率且提高数值稳定性。
- 能推导在线更新中的最大值与分母重标定公式。
- 能说明安全实现的前三个逻辑阶段如何融合为两遍。
- 能区分算法遍历数、kernel 数与真实内存流量。
前置与衔接
对向量 ,Softmax 定义为
直接实现需要先得到分母,才能写出每个 。
安全 Softmax 还要先求全局最大值,避免指数溢出,于是看起来需要三次遍历。
在线 Softmax 的关键是:最大值变化时,把此前累计的分母转换到新的指数基准,从而在同一遍流式扫描中同时维护最大值与分母。
核心讲解
1. 朴素 Softmax 的两阶段依赖
以视频中的
为例。
朴素 Softmax 先计算指数和,再用同一分母归一化各元素,因此至少包含归约与归一化两个阶段。
原视频 · 00:20 ↗第一遍计算
第二遍再写出
在不知道 之前,不能得到最终归一化结果。
因此朴素的独立数组实现至少有“归约 + 写出”两个逻辑阶段。
2. 为什么直接取指数可能不安全
若某个 很大, 可能超出浮点数可表示范围。
即使最终比例本应有限,中间指数也可能先变成正无穷,导致不定比值或 NaN。
Softmax 对整体平移不变:
所以可以选择 。
3. 安全 Softmax 把最大指数压到 1
视频示例的全局最大值为 5。
减去它后得到
安全 Softmax 将 [1,2,3,5] 减去全局最大值 5,得到 [-4,-3,-2,0],结果不变但指数不再溢出。
原视频 · 01:20 ↗此时最大的指数是 ,其他指数都位于 ,显著降低上溢风险。
画面也证实了 ASR 中漏掉的负号:移位向量是“[-4,-3,-2,0]”,不是“[4,-3,-2,0]”。
4. 朴素安全实现为何是三遍
安全 Softmax 可按以下方式实现:
- 第一遍:求 ;
- 第二遍:求 ;
- 第三遍:写出 。
朴素安全实现依次求全局最大值、求移位指数和、写出归一化结果,共三次逻辑遍历。
原视频 · 02:00 ↗第二遍依赖第一遍的 ,第三遍又依赖第二遍的 。
若把三步机械实现为三次读取输入,就会产生三遍数据访问。
5. 在线状态只需维护两个量
在线 Softmax 处理到前 个元素时维护
和
在线 Softmax 为每个分块维护局部最大值 与相对该最大值的指数和 。
原视频 · 02:40 ↗是当前前缀最大值。
不是原始指数和,而是所有已见元素相对于当前最大值 的指数和。
这两个量足以表示已经扫描部分的安全归一化信息。
6. 新元素到来时怎样更新
读入 后,先更新
旧分母 是以 为基准。
要换成新基准 ,必须乘 ,于是
若新元素没有刷新最大值,旧分母系数为 1。
若新元素更大,旧的所有指数项会统一缩小到新基准。
7. 两个分块如何合并
设分块 a、b 分别维护
先取
再把两个局部分母都变换到基准 :
合并分块时先取新最大值,再用 重标定各局部分母,保证它们处于同一指数基准。
原视频 · 03:00 ↗这就是板书中用两个局部向量信息更新全局 max 与 sum 的公式。
它允许树形归约、线程块局部归约或流式扫描,不要求先单独遍历完整向量求最大值。
8. 用 [1,2] 和 [3,5] 验证
左块:
右块:
合并最大值为 ,所以
这正好等于直接对“[-4,-3,-2,0]”求指数和。
9. 为什么从三遍降到两遍
在线更新把“求最大值”和“求相对指数和”融合到第一次扫描:
第二遍再根据最终 写出
最大值与分母可在一次流式遍历中共同更新,第二遍只负责归一化输出,从三遍降为两遍。
原视频 · 03:20 ↗所以独立 Softmax 数组实现从三次逻辑遍历变成两次。
少掉的是单独求最大值的那一遍,不是消除最终归一化。
10. 数值稳定性没有被牺牲
在线状态始终以“当前已见最大值”为指数基准。
新最大值只会不减,所有指数的自变量都满足
因此不会为少一次遍历而恢复到直接计算巨大指数的不安全做法。
有限精度下,不同归约树仍会产生微小舍入差异,但算法维持了 max-shift 的稳定原则。
11. 循环数不等于 kernel 数
视频用 for 循环与 HBM 到 SRAM 搬运解释收益。
这是很好的数据流直觉,但工程上还要区分:
- 源代码循环次数;
- GPU kernel 启动次数;
- 一个 kernel 内的线程归约;
- 编译器是否已经融合;
- 数据是否真正离开 cache 或寄存器。
若安全 Softmax 已由高度融合 kernel 实现,不能机械地把“三个公式阶段”换算成“三次完整 HBM 往返”。
在线公式提供可融合性,真实性能仍需测量。
12. 与 FlashAttention 的关系
FlashAttention 不把完整分数行存到 HBM 后再跑独立 Softmax。
它在遍历 score tiles 时维护每行的运行最大值、分母和输出累加器。
新 tile 改变最大值时,不仅要重标定 ,还要重标定此前的输出部分和。
因此在线 Softmax 是 FlashAttention 分块等价性的核心组件,但实际 kernel 比本课的独立向量两遍算法更进一步融合了 。
13. 论文与版本边界
视频把在线 Softmax 追溯到 NVIDIA 研究者 2018 年的工作,并说明它被 FlashAttention 使用。
本课聚焦 max/sum 合并公式,不把所有现代 kernel 的具体访存次数都等同于这份板书。
FlashAttention 不同版本、硬件后端和编译器可能采用不同 tile、warp 归约与流水方式,但要保持精确结果,都必须等价维护相应归一化状态。
跟练与练习
原视频练习
- 从 00:20 开始,比较朴素 Softmax 的两遍依赖
- 从 01:01 开始,理解分子分母同时乘同一缩放因子
- [从 01:10 开始,把示例移位为 [-4,-3,-2,0]](https://www.bilibili.com/video/BV1i6cnzzE2x/?p=1&t=70)
- 从 01:40 开始,数清安全 Softmax 的三次遍历
- 从 02:12 开始,观察在线 Softmax 融合前两次归约
- 从 02:41 开始,合并两个分块的 max 与 sum
- 从 03:21 开始,确认第二遍仍需最终归一化
- 从 03:26 开始,连接算子融合与内存搬运
编者练习
已扫描状态为 ,新元素为 5。更新后的 与 是什么?
查看参考答案
。旧分母必须从基准 3 转到基准 5:
不能直接写成 ,否则旧项和新项不在同一指数基准。
编者练习 2
为什么在线 Softmax 的独立数组版本仍通常需要第二遍?
查看参考答案
第一次扫描结束前,最终最大值 与分母 尚未确定,此前元素的最终概率无法永久写出。第二遍使用最终状态计算 。若与后续算子融合,可改为维护并重标定输出累加器,但那属于更大的融合算法。
常见误区
- 误区:安全 Softmax 改变了概率分布。纠正:分子分母乘同一 ,结果完全相同。
- 误区:在线 Softmax 不再减最大值。纠正:它持续减去运行最大值,并在最大值变化时重标定旧和。
- 误区:合并局部分母时直接做 。纠正:两者基准不同,必须乘相应缩放因子。
- 误区:在线 Softmax 一遍就能写出完整概率数组。纠正:独立输出通常仍需第二遍归一化。
- 误区:三次逻辑循环必然等于三次独立 HBM 读写。纠正:真实 kernel、cache 与编译融合会改变物理流量。
- 误区:FlashAttention 只维护 max 与 sum。纠正:它还维护并重标定输出累加器。
本课小结
- 朴素安全 Softmax 逻辑上依次求最大值、分母与归一化结果。
- 在线 Softmax 用状态 在一次扫描中共同更新最大值与安全分母。
- 最大值刷新时,旧分母必须乘 。
- 独立 Softmax 因此从三遍降到两遍,同时保持 max-shift 数值稳定性。
- 访存收益取决于真实融合与硬件;FlashAttention 还把在线状态扩展到输出累加。
主题讲解 · 03:25
FlashAttention 如何跳过因果掩码的上三角块
学习目标
- 能从因果掩码推导上三角 Softmax 权重为零。
- 能区分完全无效 tile、完全有效 tile 与跨越对角线的边界 tile。
- 能解释为何跳过无效 K/V tile 同时减少计算与数据搬运。
- 能说明局部分数 tile 为何仍需在线 Softmax 修正累加。
- 能区分 Decoder-only 架构与单 token decode 阶段。
前置与衔接
Decoder-only Transformer 要保证位置 不能看到未来位置 。
传统写法先计算完整分数矩阵,再在上三角加 。
FlashAttention 按 tile 计算,不必真的先算出那些注定被掩码为零的完整块。
核心判断是:
核心讲解
1. 因果注意力的掩码定义
忽略多头与缩放因子,原始分数为
因果掩码 定义为
因果注意力把位置 的未来分数设为 ,保证第 个 query 不读取未来 token。
原视频 · 00:20 ↗掩码后分数为
第 行只有当前及历史 key 是有限值。
2. 为什么负无穷对应零贡献
逐行 Softmax 为
对未来位置 ,
所以
逐行 Softmax 后,上三角的 位置权重严格为 0,不会对加权和 作贡献。
原视频 · 01:00 ↗随后输出
中,这些项都是 ,不会贡献任何值。
3. 从元素零贡献提升到 tile 零贡献
把 query 行与 key 列都分块。
若某个块中的每一对 都满足 ,该块所有掩码分数都是 。
整个 Softmax 权重 tile 都为零,进而
所以不必先计算分数再发现它为零。
运行时可以根据 tile 的位置关系预先判定并跳过。
4. 三类因果 tile
对齐分块后,可以把 tile 分为三类:
- 对角线左下方:全部满足 ,整块有效;
- 严格右上方:全部满足 ,整块无效;
- 穿过主对角线:同时包含历史与未来位置,必须在块内应用因果掩码。
只有第二类可以整块省略。
“省略上三角计算”不等于对角线边界块里完全没有 mask 判断。
5. 一个严格的区间判定
作为编者补充,设 query tile 行区间为半开区间
key tile 列区间为
若
则最小 key 下标也大于最大 query 下标,整块完全位于未来,可直接跳过。
若区间跨过对角线,就不能用整块零替代。
6. FlashAttention 的 tile 数据路径
视频以一个 query tile 为例。
线程块把 与所需的 搬入片上 SRAM,计算局部分数
FlashAttention 以 query tile 为工作单元,把需要的 Q/K 分块搬入片上 SRAM 计算局部分数。
原视频 · 01:40 ↗这里的 tile 可以包含多个 token,尺寸由片上容量、数据类型和 kernel 配置共同决定。
7. 局部分数不写回 HBM
在片上立即参与 Softmax 状态更新与 加权。
局部注意力分数块 在片上参与 Softmax 与后续乘法,使用完即丢弃,不写回 HBM。
原视频 · 02:00 ↗使用完成后即可丢弃,不必物化到 HBM。
这与“跳过上三角”是两项相关但不同的优化:
- 不物化 A:避免完整分数矩阵的 HBM 存储与往返;
- 跳过无效 tile:连分数乘法与对应 K/V 搬运都不做。
8. 为什么一个局部 O 不是最终输出
当前块得到的部分结果类似
但 还可能关注 中的合法位置。
每个分数块只覆盖部分 key,所得 必须结合在线 Softmax 状态修正并累加到输出 tile。
原视频 · 02:40 ↗因此它必须与后续 K/V tile 的贡献累加。
更重要的是,每个分数块只有局部最大值和分母。
新块改变运行最大值时,旧输出累加器也要按在线 Softmax 比例重标定,才能与完整一行 Softmax 等价。
9. 为什么可以不搬未来 K/V tile
考虑最早的一组 query,例如 ,与很靠后的 。
若二者对应的整个分数块都位于因果上三角,那么该块最终全为零。
若某个 K/V tile 完全位于当前 query tile 的因果未来区域,其分数全为 ,可直接不搬运、不计算。
原视频 · 03:00 ↗此时可以同时省去:
- 从 HBM 加载 ;
- 从 HBM 加载对应 ;
- 计算 ;
- 对该零块执行 Softmax 更新;
- 计算零权重与 的乘法。
收益不只是少做乘法,也包括少搬数据。
10. 跳过不会破坏 Softmax 分母
有人会担心:不计算这些项,Softmax 分母是否会改变?
不会。
被掩码项的指数本来就是
把零项从求和中删除不改变分母,也不改变最大值,因为合法行至少包含当前位置本身。
因此跳过是代数上的精确消元,不是近似稀疏化。
11. Prefill 与单 token decode 要分开
视频 ASR 把 “Decoder-only Transformer” 识别得不清楚,不能据此把课程理解成单 token decode。
在处理完整 causal 序列的训练或 Prefill 中, 和 都覆盖许多位置,确实存在大块上三角区域。
而标准自回归单 token decode 只计算最新 query,它可以关注所有已有 key,通常没有同样的完整上三角矩阵可跳。
因此该优化的显著场景主要是多 query 的 causal attention。
12. 负无穷在实现中不一定真的写入
数学上用 表示 mask 最清楚。
kernel 中可能:
- 对边界 tile 生成布尔 mask;
- 用足够小的有限值实现;
- 直接让无效 lane 不参与归约;
- 对完全无效 tile 在调度层跳过。
只要最终等价于无效位置的 Softmax 权重为零,具体表示可以不同。
13. 版本边界
视频以 FlashAttention v2 讲解外 Q tile 与因果上三角跳过。
这一“完全掩码 tile 不加载、不计算”的原则并不应理解为只有 v2 才能采用。
FlashAttention 不同版本、不同后端以及 PyTorch SDPA/Triton 等实现,可能有不同分块大小、对角线处理和调度策略。
判断某版本是否实际跳过哪些 tile,应查看目标 kernel 与运行配置。
跟练与练习
原视频练习
编者练习
query tile 覆盖行 ,key tile 覆盖列 。按 为未来位置,这个 tile 能否整块跳过?
查看参考答案
可以。最小 key 下标 6 仍大于最大 query 下标 5,所以四个位置组合都位于因果未来区域,Softmax 权重全为零。
编者练习 2
query tile 覆盖 ,key tile 覆盖 ,能否整块跳过?
查看参考答案
不能。对 ,位置合法;而 、 等位置被掩码。该 tile 跨越对角线,必须计算合法部分并在块内应用 mask。
常见误区
- 误区:先计算上三角,再把结果清零。纠正:完全掩码 tile 可在加载和乘法前就跳过。
- 误区:所有碰到上三角的 tile 都能整块跳过。纠正:跨对角线的边界 tile 仍含合法元素。
- 误区:不搬 K 就仍要搬对应 V。纠正:该权重块全零时,K 与 V 的整个贡献路径都可省去。
- 误区:跳过掩码项会改变 Softmax 分母。纠正:这些项的指数本来就是零。
- 误区:局部 已是最终输出。纠正:还要跨 K/V tiles 做在线修正与累加。
- 误区:Decoder-only 就等于单 token decode。纠正:架构名称与推理阶段不同,上三角主要出现在多 query causal 计算。
本课小结
- 因果掩码令 的分数为 ,Softmax 权重为零。
- 完全位于上三角的 tile 对输出没有贡献,可以不加载 K/V、不计算分数。
- 跨越主对角线的 tile 仍需块内 mask,不能整体删除。
- FlashAttention 同时避免把局部分数写回 HBM,并用在线 Softmax 修正累加部分输出。
- 该跳过是精确消元;具体 tile 调度和实现细节随版本与后端变化。
主题讲解 · 02:55
FlashAttention-2 为什么把 Q 放到外层循环
学习目标
- 能画出外层 Q、内层 K/V 的 tile 循环。
- 能说明哪些张量常驻 HBM,哪些状态只需片上暂存。
- 能解释同一个输出 tile 为何要跨多个 K/V tile 修正累加。
- 能比较 FlashAttention v1 与 v2 循环顺序的 Q/O 搬运差异。
- 能区分视频聚焦的访存理由与 FlashAttention-2 的其他优化。
前置与衔接
FlashAttention 计算的数学目标仍是
它不把完整注意力分数矩阵写回 HBM,而是让局部分数 tile 在片上被消费。
当 key 轴也被切块时,一个 query tile 的最终输出依赖所有合法 K/V tiles。
循环顺序决定:究竟让当前 query/output tile 留在片上,还是让当前 K/V tile 留在片上。
FlashAttention-2 选择外层 Q、内层 K/V,视频用“Q/O 定、K/V 动”概括其访存收益。
核心讲解
1. 先区分持久张量与临时张量
完整 Q、K、V 与最终输出 O 通常位于 HBM。
完整 Q、K、V 与最终 O 位于 HBM;计算围绕一个 query/output tile 在 HBM 与 SRAM 间组织。
原视频 · 00:20 ↗分块计算会把当前需要的小块搬到寄存器或共享内存等片上存储。
完整分数矩阵
则不需要持久化到 HBM。
2. 分数 tile 使用后即可丢弃
对 query block 与 key block ,先计算
局部 参与在线 Softmax,再与 相乘。
局部分数 与对应 V 块在片上使用,分数 tile 无需物化回 HBM。
原视频 · 00:40 ↗一旦它对在线统计量和输出累加器的贡献已经吸收,原始分数 tile 就可以释放。
因此循环顺序主要影响 Q、K、V、O 及在线状态怎样搬运,而不是 A 如何写回。
3. 固定一个 query tile
视频选 作为例子。
把它搬入片上存储后,与第一个 K block 计算
再读取 ,形成第一份局部输出贡献。
之后依次处理
固定 后,依次与 计算 ,共同覆盖完整 key 轴。
原视频 · 01:40 ↗内层循环结束时,当前 query 行已经看过全部合法 key。
4. 为什么不能简单累加局部 Softmax
设第一个分数块的局部最大值和分母为 ,第二块为 。
若分别对两个块归一化再直接相加,会让每块权重各自和为 1,结果不等于完整行 Softmax。
必须维护运行状态
旧输出部分和也要按相同基准重标定。
5. 每个 K/V block 都贡献到同一个 O tile
对 ,第一块产生 ,第二块产生 。
、 等部分结果经在线 Softmax 修正后累加到同一个 。
原视频 · 02:00 ↗更准确的在线更新形式可写为
其中 是尚未除以最终分母的加权和。
等所有 K/V blocks 处理完,再除以最终 得到 。
6. 为什么 O2 必须跨内层循环保留
不是一次 就完成的。
它与在线最大值 、分母 一起构成跨 K/V blocks 的运行状态。
如果每处理一个 K/V block 都把这组状态写回 HBM,下一轮又读回来,会产生重复访存。
更好的做法是让负责 的线程块在内层循环期间保留这些状态。
7. 外 Q 内 K/V 的片上驻留
FlashAttention-2 的抽象循环可以写成:
- 选择一个 query tile ;
- 把 放入片上存储;
- 初始化该 tile 的 ;
- 依次流过所有合法 ;
- 完成后一次写回最终 。
一个线程块可让 与运行中的 驻留 SRAM,只流式替换 K/V tile。
原视频 · 02:20 ↗这就是“Q/O 定,K/V 动”。
这里的“SRAM”是视频的简化表达;具体 kernel 可能把不同状态分配在 shared memory 与 registers。
8. FlashAttention v1 的相反循环
视频指出 v1 的高层伪代码先遍历 K/V blocks,再在内层遍历 Q blocks。
在这种组织下,一个 K/V tile 可以被多个 query tiles 复用。
但对某个 而言,每个外层 K/V block 到来时,都可能需要:
- 从 HBM 重新读入 ;
- 读回此前的 ;
- 更新后再写回。
于是 Q 与输出状态可能随 K/V 外层循环反复搬运。
9. v2 交换循环后的核心收益
交换为外层 query 后,一个线程块从头到尾负责一个 tile。
FlashAttention-2 采用外层遍历 query tiles、内层遍历 key/value tiles 的组织,让当前 Q/O tile 固定、K/V 块移动。
原视频 · 02:40 ↗它只需在开始时读取 ,在结束时写回一次最终 。
在线 Softmax 状态也可在片上存活整个内层循环。
因此减少了 Q、O 以及 的 HBM 往返。
10. 交换循环不改变数学结果
对固定 query block,完整输出是对所有 key blocks 的贡献归约。
只要在线 Softmax 的重标定正确,按 的顺序逐块合并,与一次计算完整 key 轴等价。
循环重排改变数据驻留与并行任务划分,不改变最终
有限精度下归约顺序可能带来小幅舍入差异,但不是算法近似。
11. K/V 重复读取是否成了新问题
外层遍历多个 Q tiles 时,同一个 K/V tile 会被不同线程块读取。
所以 v2 不是让所有张量都只读一次,而是选择减少更昂贵或更频繁的 Q/O 状态往返,并换取更好的 query 维并行。
真实性能取决于:
- Q/K/V block 大小;
- 片上容量;
- cache 命中;
- 序列长度与 head dimension;
- batch/head 并行度;
- causal mask 能跳过多少 K/V blocks。
循环顺序是 I/O 权衡,不是无条件消除所有重复加载。
12. 为什么并行度也更自然
外层 query tiles 彼此输出独立,可以由不同线程块并行处理。
每个线程块只需维护自己的 。
这为较长序列提供 query 维的并行任务。
FlashAttention-2 的论文与实现还包含减少非矩阵乘 FLOPs、改善线程块/warp 工作划分等优化。
视频本课只聚焦循环交换与 Q/O 搬运,不应把 v2 的全部收益都归因于这一点。
13. 与因果上三角跳过的结合
在 causal attention 中,固定 后,内层只需遍历合法 K/V blocks。
严格位于未来的 blocks 可以直接跳过;对角线边界块执行块内 mask。
因此外 Q 内 K 的结构也便于为每个 query tile 确定自己的 key 范围。
但具体 kernel 如何计算边界、是否使用 split-K 或不同调度,随实现变化。
14. 版本边界
本课按视频解释 FlashAttention v1/v2 论文级高层循环差异。
现代 FlashAttention、框架集成和不同 GPU 架构可能采用多套 kernel;短序列、decode、varlen、causal、GQA 等路径不一定共享同一伪代码。
判断当前软件版本的真实数据驻留和循环次序,应查看对应 kernel、编译参数与 profiler,而不是只凭版本名。
跟练与练习
原视频练习
编者练习
有 3 个 Q tiles 和 4 个 K/V tiles。若 v1 高层循环外 K 内 Q,而每次更新都需从 HBM 读取并写回 O tile,那么每个 O tile 会被更新几次?交换循环后理论上可在何时写回?
查看参考答案
外 K 内 Q 时,每个 O tile 会随 4 个 K/V tiles 更新 4 次,若状态无法跨外层保持,就会多次读写。外 Q 内 K 时,一个线程块可让该 O tile 跨 4 次内层迭代驻留,处理完成后一次写回最终结果。
编者练习 2
为什么 不能直接当作 Softmax attention 的输出?
查看参考答案
是原始分数而非完整行归一化权重。即使各块先做局部 Softmax,也缺少跨块共同最大值与分母。必须用在线 Softmax 维护并重标定 max、sum 与输出累加器。
常见误区
- 误区:FlashAttention-2 把完整 A 留在 SRAM。纠正:只暂存当前分数 tile,消费后丢弃。
- 误区:O2 由一个 K/V tile 就能算完。纠正:它需要所有合法 key blocks 的修正累加。
- 误区:外 Q 内 K 让 K/V 只读一次。纠正:不同 Q 线程块仍可能重复读取 K/V。
- 误区:交换循环改变了注意力定义。纠正:正确在线归约后,数学结果仍是精确 attention。
- 误区:v2 的全部加速只来自循环交换。纠正:还有并行划分与非矩阵乘开销等改进。
- 误区:所有当前 v2 kernel 都严格对应同一板书。纠正:路径随硬件、shape、mask 与软件版本变化。
本课小结
- FlashAttention-2 让一个线程块围绕固定 Q/O tile 工作,内层流过 K/V tiles。
- 局部分数 tile 不写回 HBM,在线 max、sum 与输出累加器跨内层循环保留。
- 与 v1 的外 K 内 Q 相比,该循环减少 Q/O 状态的反复 HBM 搬运。
- 循环重排保持精确注意力,但仍有 K/V 读取与硬件调度权衡。
- 视频解释的是 v2 的一个关键访存动机,不覆盖所有现代 kernel 细节。
主题讲解 · 02:32
物化注意力分数矩阵为何如此昂贵
学习目标
- 能从 Q、K 的 shape 推导 分数矩阵。
- 能区分平方中间存储与稠密注意力的平方计算量。
- 能解释 FlashAttention 如何避免把完整 A 写回 HBM。
- 能说明 关于 为线性的成立条件。
- 能区分前向、反向与整个模型显存的不同统计口径。
前置与衔接
FlashAttention 经常被概括为把注意力显存从平方级降到线性级。
这句话容易被误解成“注意力不再计算 个 query-key 关系”。
视频抓住了更准确的瓶颈:
数学输出仍是精确的 dense attention,主要改变的是内存复杂度与数据搬运。
核心讲解
1. 先固定 N 与 d 的含义
对单个 attention head,设:
- :序列 token 数;
- :该 head 的维度。
则
单个 attention head 中, 与 的 shape 为 ,视频以 、 说明两轴量级差异。
原视频 · 00:40 ↗视频用长上下文 、head dimension 作量级示例。
具体模型的 可以不同,但长序列时通常关注 的情形。
2. 为什么 会变成
转置后
因此
的 Q 与 的 相乘,输出分数矩阵沿 query、key 两轴扩展为 。
原视频 · 01:00 ↗第一根 N 轴枚举 query token,第二根 N 轴枚举 key token。
每个元素
表示一对 token 的注意力分数。
3. 平方存储从哪里来
若传统实现把 A 写回 HBM,它需要保存
个元素。
传统实现若把 的 分数矩阵写回 HBM,会引入关于序列长度的 存储。
原视频 · 00:20 ↗而单个 Q、K、V 或 O 只包含
个元素。
两者比值为
当 时,A 很快成为更大的中间张量。
4. 用 N=10000、d=128 感受量级
单头 Q 的元素数为
A 的元素数为
A 约为单个 Q 的 78 倍。
若仅为说明量级,按 FP16/BF16 每元素 2 字节,单头 A 约 200 MB。
实际总量还会乘 batch、head 数,并受 dtype、是否保存 logits/probabilities、mask 与框架实现影响。
这个数字只是 shape 估算,不是任意模型运行时的固定显存值。
5. 不好不只在“占着显存”
把 A 写回 HBM 后,后续 Softmax 和 还要把它读回来。
因此物化 A 同时带来:
- 大型中间张量的峰值显存;
- 写 A 的 HBM 流量;
- 读 A 或 P 的 HBM 流量;
- 额外 kernel 边界与同步机会;
- 长序列下更强的带宽压力。
许多 attention kernel 在实际硬件上受内存 I/O 限制,而不只是算术吞吐限制。
6. FlashAttention 怎样切 tile
FlashAttention 把 Q 沿 query token 行切成 tiles,把 沿 key token 列切成 tiles。
FlashAttention 把 Q 沿 token 行切块、把 沿 key 列切块,每个 tile 通常包含多个完整 token 向量。
原视频 · 01:20 ↗在视频的简化画法里,每个 token 向量的完整 维保留在 tile 中。
一个 的 Q tile 与 的 tile 相乘,得到
根据片上容量和 kernel 设计选择。
7. 为什么不沿 d 随便切断
每个注意力分数需要完整收缩
若沿 切分,也不是数学上不可能,但必须再对各段部分内积做归约。
视频强调“token 完整性”,是在当前 tile 解释中把完整 head dimension 放进计算块,避免再引入 contraction 轴的跨块合并。
不能把它理解为任何 kernel 都绝不会在硬件微操作层拆分 。
8. A tile 只在片上短暂存在
局部分数 产生后,立即用于更新:
- 每行运行最大值;
- 每行 Softmax 分母;
- 对 V 的输出加权和。
每个小分数块只在 SRAM 中参与在线 Softmax 与 加权,消费后丢弃,不形成 HBM 中的完整 张量。
原视频 · 02:00 ↗它不写回 HBM。
随后下一个 K/V tile 覆盖另一段 key 轴,在线 Softmax 公式把新旧局部状态精确合并。
所有 key blocks 处理完后,只写回最终 O tile。
9. 为什么结果仍等于完整 Softmax
分块只改变求值顺序。
对每个 query 行,FlashAttention 维护全局等价状态
以及未归一化输出和
每来一个新 tile,若最大值变化,就重标定旧 与 。
最终
与完整矩阵 Softmax 相同。
10. 到底指什么
不物化 A 后,主要持久 attention 张量 Q、K、V、O 都是 量级。
不物化 A 后,主要持久张量 Q、K、V、O 均为 量级;当 d 视为固定时,额外存储关于 N 为线性。
原视频 · 02:20 ↗所以 attention 的额外持久存储可写作
当 对所研究的序列长度 N 视为固定模型参数时,它关于 N 是线性的。
更严谨的说法是“从 中间存储降到 ”,而不是无条件把两个变量都省略成 。
11. 计算复杂度仍然是平方级
Dense attention 仍需覆盖所有合法 query-key 对。
的算术量级仍是
也有同阶计算。
FlashAttention 的主要贡献是 I/O-aware:减少 HBM 流量与中间存储,从而让相同 dense attention 计算在硬件上更高效。
它没有把稠密注意力变成线性 attention 算法。
12. Causal attention 的常数会变化
因果掩码只保留下三角,合法 query-key 对约为
这能把常数约减半,并允许跳过完全上三角 tiles。
但渐近计算量仍是 。
因此“跳过上三角”和“不物化完整 A”应分别理解:前者减少无效计算,后者改变中间存储与 I/O。
13. 训练反向为何也能省存储
训练需要反向传播通过 Softmax。
传统自动微分常保存大型中间激活。
FlashAttention 的反向可以利用前向保存的较小统计量和输出,重新计算局部分数与概率,而不是保存完整 A/P。
这用额外重计算换取更低内存。
具体保存哪些统计量、dropout RNG 如何恢复,会随算法版本与框架实现变化。
14. 不要把它推广成“整个模型显存线性”
模型总显存还包括:
- 参数、梯度和优化器状态;
- MLP 与归一化激活;
- KV Cache;
- logits、loss 与其他 workspace;
- 框架内存池和碎片。
FlashAttention 消除的是 attention 中最关键的 中间物化。
它不保证整个训练或推理进程的全部显存只剩 。
15. 版本边界
视频用“FlashAttention 之前的 注意力”概括传统 materialize-then-softmax 路径。
现代框架即使不显式标为 FlashAttention,也可能通过 memory-efficient SDPA、Triton kernel 或其他融合实现避免完整 A。
FlashAttention v1/v2/v3 与不同硬件后端的 tile、并行和反向策略也不同。
判断当前运行是否真的物化 中间量,应检查所选 backend、shape 支持条件和 profiler,而不是只看 API 名称。
跟练与练习
原视频练习
编者练习
时,A 与单个 Q 各有多少元素?A 是 Q 的多少倍?
查看参考答案
有 个元素;Q 有 个元素。二者比值为 ,所以 A 是单个 Q 的 64 倍。
编者练习 2
FlashAttention 不保存完整 A,为什么计算复杂度仍不是 ?
查看参考答案
不保存不等于不计算。每个 query 仍需与所有合法 key 做内积,分块只让这些分数在片上短暂产生并立即消费。因此 dense attention 的 query-key 配对数仍为平方级,算术量约 。
常见误区
- 误区: 本身就是 。纠正:单头 为 ,乘积 才是 。
- 误区:FlashAttention 把计算复杂度降到线性。纠正:它主要把中间存储降到 ,dense 计算仍为 。
- 误区:A 在 SRAM 中完整保存。纠正:只保存当前小 tile,使用后丢弃。
- 误区:不写 A 就得近似 Softmax。纠正:在线 max、sum 与输出重标定保持精确等价。
- 误区: 意味着 也可任意增长而仍关于所有变量线性。纠正:称关于 线性时把模型维度 视为固定。
- 误区:FlashAttention 让整个模型显存都变成 。纠正:它针对 attention 的平方中间量,其他状态仍存在。
本课小结
- 单头 为 ,而分数矩阵 为 。
- 传统物化 会带来 存储及大规模 HBM 读写。
- FlashAttention 分块产生 A tile,在片上完成在线 Softmax 与 V 加权后立即丢弃。
- 主要持久 attention 张量降为 ,当 固定时关于 线性。
- dense attention 的计算仍为 ;版本与框架后端决定实际是否物化中间矩阵。
主题讲解 · 01:47
长序列下注意力显存差距如何被放大
学习目标
- 能区分序列长度 N 与单头维度 d。
- 能从 shape 推导传统 attention 的 中间存储。
- 能计算 与 的元素数比值。
- 能解释为什么短序列下差距不明显、长序列下差距剧增。
- 能准确说明 FlashAttention 的“线性显存”指存储而非算术量。
- 能辨别单个中间张量、attention 子层和整个模型显存三种口径。
前置与衔接
上一课已经解释了完整注意力分数矩阵为何是 。
本课进一步问一个更量化的问题:
答案不是固定的“慢若干倍”。
差距随 N 增大而继续放大。
要看清这个趋势,先固定单个 attention head 的记号:
其中:
- N 是 token 数;
- d 是该 head 的维度;
- 比较长上下文时,常见情形是 。
总览板书把传统方法的 注意力中间量与 FlashAttention 的 持久张量并列,比较关于序列长度 N 的平方与线性存储。
原视频 · 00:00 ↗核心讲解
1. 两种增长来自不同 shape
传统实现先计算
因为
所以
若把 A 或由它得到的 Softmax 概率写回 HBM,单头就是 个元素。
视频在 00:14 开始把这一路径称为平方存储。
传统实现若物化 A 或 Softmax(A),query 与 key 两根轴都随 N 增长,因此该中间量包含 个元素。
原视频 · 00:40 ↗FlashAttention 仍计算精确的 dense attention。
不同点是它按块消费分数、执行在线 Softmax 并累加输出,不让完整 A 成为 HBM 中的持久中间量。
因此主要持久张量仍是 量级。
2. 为什么短序列时看不出危险
视频在 00:24 先画了 N 很小时的情况。
若 N 与 d 同量级, 和 的数值不会相差太远。
例如 N=10、d=8:
只看这个规模,完整分数矩阵似乎“不算大”。
当 N 较小时, 产生的 矩阵在图形上并不比 张量大很多,平方增长的危险尚不直观。
原视频 · 00:20 ↗这正是容易形成误判的地方:
- d 通常由模型架构固定;
- N 会随着上下文长度扩展;
- 平方项同时在 query 与 key 两根轴上增长。
3. 比值不是常数,而是 N/d
把 A 与一个 张量比较:
因此,只要 d 固定,N 每扩大一倍,这个相对倍数也扩大一倍。
若把 Q、K、V、O 四个 张量合计作为参照,则比值为
这两个比值口径不同:
- :A 相对单个 张量;
- :A 相对 Q、K、V、O 元素数总和。
不能把二者混用。
4. 用视频的长序列示例计算
长序列示例中 ;板书以 、 强调 与 的量级差会迅速拉大。
原视频 · 01:00 ↗一个 张量包含
个元素。
完整 A 则包含
个元素。
A 相对单个 张量大
倍。
相对四个 Q、K、V、O 的元素数总和,仍约大
倍。
若只作元素字节的理想化估算,FP16 的 A 单头就需要约
这个数字只用于说明平方项规模;真实实现还受 batch、head 数、mask、重计算、融合策略和数据类型影响。
5. FlashAttention 省掉的到底是什么
视频在 00:50 明确指出 FlashAttention 不写回 A。
在 01:27 又用长序列图重申这一点。
FlashAttention 不把完整 注意力矩阵写回 HBM,主要持久张量保持在 量级;这里说的是存储而非稠密 attention 的算术量。
原视频 · 01:20 ↗所以“线性”更准确地说是:
它不意味着:
- 稠密 attention 的 query-key 配对数变成 ;
- 整个模型训练显存只剩 ;
- 所有实现的常数项都相同;
- 长序列没有计算代价。
标准 dense attention 的核心算术量仍约为 。
跟练与练习
编者练习
把平方项换算成倍数 单个 head 取 N=32768、d=128。
- A 有多少个元素?
- 一个 张量有多少个元素?
- A 是单个 张量的多少倍?
- 若两者均为 FP16,A 理想化占多少 GiB?
查看参考答案
比值为
FP16 按每元素 2 bytes 估算:
这里只计算一个单头 A,不代表完整训练显存。
快速判断
判断下列说法:
- “N 固定时,增大 d 会让 A 本身变大。”
- “d 固定时,N 翻倍会让 A 元素数变为四倍。”
- “FlashAttention 不存完整 A,所以不再计算全部 dense attention 关系。”
答案依次是:错、对、错。
常见误区
误区 1:线性比平方永远固定快某个倍数
错。
相对倍数含有 ,会随 N 变化。
误区 2:只写 ,忽略 d
更严格的 shape 口径是 。
只有把 d 视为固定模型常数时,才简称“关于 N 线性”。
误区 3:把内存复杂度当成计算复杂度
FlashAttention 的主要贡献是 IO-aware 分块与不物化平方中间量。
标准稠密 attention 的算术关系仍是平方级。
误区 4:把 A 的大小当成全部显存
完整训练还包括参数、优化器状态、其他层激活、梯度、临时 workspace 等。
本课只比较 attention 路径中的关键中间量。
误区 5:看到 FP16 估算就认为实现一定分配同样字节
真实 kernel 可能重计算、融合、使用不同累加精度,也可能受布局与对齐影响。
元素数分析用于判断渐近主项,不替代 profiler。
本课小结
- 传统物化注意力矩阵需要 个元素。
- FlashAttention 主要持久张量处于 量级。
- A 相对单个 张量的元素数比值是 。
- N 很小时差距不明显,N 远大于 d 时差距快速放大。
- N=100000、d=128 时,A 相对单个 张量约大 781 倍。
- “线性显存”是关于 N 的存储结论,不是 dense attention 的线性计算结论。
主题讲解 · 03:07
在线 Softmax 与 FlashAttention 的状态量
学习目标
- 能说明在线 Softmax 每行维护的两个核心状态。
- 能说明 FlashAttention 前向为何还需要输出累加状态。
- 能推导跨块最大值、分母和输出分子的重标定公式。
- 能区分未归一化输出累加器与已归一化输出的两种记法。
- 能解释“动态规划量”是教学类比,而非传统 DP 表格。
- 能说明不同 FlashAttention 版本与 kernel 实现的存储细节边界。
前置与衔接
普通安全 Softmax 对一行 logits 先取
再计算
概率为
在线 Softmax 的关键是:不必一次看到整行,仍能维护出同一个 m 与 。
视频在 00:06 特别说明,把这看成“动态规划”主要是帮助理解局部状态如何维护全局结果。
它并不是经典的二维 DP 表格。
核心讲解
1. 在线 Softmax 的两个状态
视频在 00:12 给出数量结论:
- 在线 Softmax:两个核心状态;
- FlashAttention 前向:三个概念状态。
在线 Softmax 每行维护运行最大值 m 与安全分母 ;FlashAttention 前向还要维护输出累加状态,概念上共有三类状态量。
原视频 · 00:20 ↗对每一行而言,在线 Softmax 的状态是:
- 运行最大值 m;
- 以 m 为指数基准的安全分母 。
这里的 是标量。
m 也是标量。
它们都是“每行一个”,不是整个 attention 矩阵只保存一个。
2. 两个块如何合并
设旧状态覆盖集合 A:
新块 B 的局部状态为
合并后的最大值是
但 与 使用了不同指数基准,不能直接相加。
先统一到新的 m:
分块安全 Softmax 先以各块局部最大值为基准计算指数和;合并时必须把旧分母重标定到新的全局最大值。
原视频 · 00:40 ↗视频从 00:45 开始解释这个分母合并。
每个缩放因子都不大于 1,因此保留了安全 Softmax 的数值稳定性。
3. 为什么局部概率不能直接拼接
视频用 与 两块演示。
第一块局部最大值是 2,第二块局部最大值是 5。
于是局部指数分别是:
和
两块 logits 各自减去局部最大值后得到局部指数权重;这些权重不能直接拼接,必须先统一指数基准。
原视频 · 01:20 ↗第二块的 与第一块的 并不代表相同的全局权重。
因为它们分别以 5 和 2 为参考。
把第一块改写到全局最大值 5 的基准,需要整体乘
这就是“局部信息维护全局信息”的核心。
4. FlashAttention 增加的第三类状态
attention 输出一行是
若 V 也分块流入 SRAM,kernel 不能只维护 m 与 。
它还要维护与 V 维度相同的输出累加状态。
片上 SRAM 容量有限, 分块流入并被替换;因此输出 O 也要随每个块在线更新,而不能等完整概率矩阵出现。
原视频 · 02:00 ↗因此,“三个量”应理解为三类每行状态:
- 标量最大值 m;
- 标量安全分母 ;
- 向量输出累加器。
第三项不是一个标量。
5. 用未归一化输出分子推导
最清楚的记法是定义块 B 的未归一化输出分子
旧累加器为
统一到新的全局最大值 m 后:
最终归一化输出是
各块未归一化输出需乘 重标定后求和,最后再除以全局分母 ;这才与整行 Softmax 加权结果等价。
原视频 · 02:40 ↗6. 若实现维护的是已归一化 O
有些讲解或实现把第三个状态直接记成已归一化输出 O。
设
则更新式要写成
这与未归一化记法完全等价。
但不能把已归一化 O 直接乘指数因子后相加,再忘掉 、。
7. “三个状态”不是所有实现字段的总数
三类状态是理解前向数学合并所需的最小概念集合。
具体 kernel 还可能保存:
- log-sum-exp;
- dropout 随机数状态;
- mask 或序列边界信息;
- tile 索引与流水线元数据;
- 反向阶段所需的辅助量。
FlashAttention-1、FlashAttention-2、FlashAttention-3 及不同框架封装的具体寄存器、SRAM 与 HBM 布局并不相同。
所以本课的“2 个/3 个”是算法状态分类,不是某个版本源码中变量名的机械计数。
跟练与练习
编者练习
合并两块 Softmax 状态 给定: 写出合并后的 m、、 与 O。
查看参考答案
全局最大值为
第一块需乘 ,第二块基准不变:
未归一化输出分子为
最后
展开可验证它等于以全局最大值 5 计算的整行 attention 输出。
自检:状态的 shape
若一个 query block 有 行,head dimension 为 d,则:
- m 的 shape 是 ;
- 的 shape 是 ;
- 的 shape 是 。
这比“总共三个标量”更准确。
常见误区
误区 1:局部 Softmax 概率可以直接拼接
错。
每块的局部最大值与局部分母不同,必须统一基准并重新归一化。
误区 2:在线 Softmax 只维护最大值
错。
最大值保证指数稳定,但还需要分母才能恢复正确概率。
误区 3:FlashAttention 的第三个状态是最终 O
不一定。
数学上既可维护未归一化 ,也可维护已归一化 O;两种更新式不同。
误区 4:“动态规划量”就是 DP 数组
本课使用的是类比。
更准确地说,这些是可合并的充分状态或归约状态。
误区 5:不同版本一定保存同样的中间量
错。
版本、精度、mask、dropout 与硬件都会改变实现细节,但最大值重标定的数学原则不变。
本课小结
- 在线 Softmax 每行维护运行最大值 m 与安全分母 。
- 跨块合并时,局部分母必须乘 后求和。
- FlashAttention 前向还要维护输出累加状态,因此概念上有三类状态。
- 未归一化输出分子与分母使用完全相同的最大值重标定因子。
- 最终输出由 得到。
- “2 个/3 个”是算法状态分类,不是某个 kernel 的全部实现变量。
主题讲解 · 03:23
用矩阵链式法则推导注意力反向梯度
学习目标
- 能从标量乘积链式法则过渡到矩阵乘积反向传播。
- 能使用 Frobenius 内积检查矩阵梯度的方向。
- 能推导 对 Q、K 的梯度。
- 能推导 对 P、V 的梯度。
- 能写出逐行 Softmax 的 Jacobian-vector product。
- 能说明缩放、mask、dropout 等实际 attention 细节应插入何处。
前置与衔接
视频试图用一句口诀快速记忆矩阵乘法反向:
这个口诀有用,但必须补上两个条件:
- 矩阵乘法不可交换,因子次序不能随意改;
- 最终公式要通过 shape 或微分内积检查。
从标量乘积 出发,对 b 求导时把上游梯度放到 b 的位置,并乘上其余因子,作为矩阵推广的直觉起点。
原视频 · 00:20 ↗核心讲解
1. 标量乘积为何可以“去掉被求导因子”
设
损失 L 通过 y 依赖 b。
则
标量乘法可交换,因此看起来像“把 b 替换成上游梯度,剩下的照乘”。
矩阵乘法没有交换律,所以推广时不能只照搬“剩下的照乘”。
2. 矩阵乘积的可靠推导工具
设
上游梯度记为
只让 B 变化:
用 Frobenius 内积定义梯度:
代入并循环移动 trace 中的因子:
因此
矩阵乘积的反向传播不能任意交换因子;对中间矩阵求梯度时要保持乘法次序,并用转置使 shape 对齐。
原视频 · 01:20 ↗视频在 01:08 开始把标量结论推广到矩阵。
板书的紧凑写法 与上式相同,因为
3. shape 检查比口诀更可靠
假设
则
候选梯度
的 shape 是
正好等于 B 的 shape。
若 shape 不对,转置或次序一定有错。
4. 注意力前向的最简计算图
本课采用视频中的简化记号:
反向从上游 dO 开始,依次得到:
这就是板书左侧所列的五条主要梯度支路。
5. 由 O=PV 推出 dP 与 dV
微分为
分别收集两支:
对 ,反向传播得到 与 ,shape 检查可快速发现转置错误。
原视频 · 02:20 ↗视频在 02:14 开始把 P、V 两支与 Q、K 两支类比。
shape 检查:若
则
与 P 相同。
而
与 V 相同。
6. 由 S=QK^T 推出 dQ 与 dK
微分为
对应梯度:
对 ,上游梯度为 dS 时有 、;K 原本带转置,因此两支公式不完全对称。
原视频 · 02:00 ↗dK 之所以看起来多一步,是因为前向中出现的是 。
可先求
再转置:
7. Softmax 反向不能套矩阵乘法口诀
的反向不是普通矩阵乘法;dS 需由 Softmax 的 Jacobian-vector product 从 dP 与 P 计算。
原视频 · 03:00 ↗对单行
其 Jacobian 为
给定上游梯度 g=dP,该行的 dS 为
矩阵逐行写成
其中 rowsum 的结果在该行广播。
8. 实际 attention 中还缺哪些因子
标准 scaled dot-product attention 常写为
若缩放因子
则
视频板书为突出矩阵链式法则,省略了这个缩放。
mask M 若为常量,不需要对 M 求梯度,但被屏蔽位置的概率与梯度处理必须符合具体实现。
训练中的 dropout、causal mask、GQA/MQA、变长序列也会改变 kernel 路径,但不改变上述局部矩阵微分原则。
9. 这套推导与 FlashAttention 的关系
这些是数学上的反向公式。
FlashAttention 的工程难点不是改变梯度定义,而是:
- 按块重计算前向需要的 S、P;
- 避免完整 中间量写回 HBM;
- 在片上内存有限条件下组织 dQ、dK、dV 的归约;
- 控制数值误差与数据搬运。
因此,“能写出五条梯度公式”不等于“已经实现 FlashAttention backward kernel”。
跟练与练习
编者练习
只用 shape 推出两条梯度 给定 其中 写出 dQ、dK,并验证 shape。
查看参考答案
shape:
shape:
两者分别与 Q、K 的 shape 相同。
进阶自检
若前向为
且 c 是常量,dQ 与 dK 都应额外乘 c。
若遗漏 c,数值梯度检查会失败。
常见误区
误区 1:矩阵乘法像标量一样可以交换
错。
转置会反转乘积次序,shape 是最直接的检查手段。
误区 2:“对谁求导就替换谁”足以覆盖所有算子
它只适合作为矩阵乘法的记忆辅助。
Softmax、mask、dropout 等需要各自的局部反向规则。
误区 3:
是对 的梯度 shape。
对 K 的梯度还需转置,得到 。
误区 4:板书省略缩放就代表实际 attention 没有缩放
本课使用的是简化计算图。
实现时要把 放回 dQ、dK。
误区 5:推导数学梯度等于解决内存问题
数学公式定义结果。
FlashAttention backward 还要决定哪些量重计算、如何分块以及如何归约。
本课小结
- 标量乘积的直觉可以推广,但矩阵因子不能交换。
- Frobenius 内积与 trace 是可靠的矩阵梯度推导工具。
- 对 ,有 、。
- 对 ,有 、。
- Softmax 反向是逐行 Jacobian-vector product,不是普通矩阵乘法。
- 实际 scaled attention 还要把缩放、mask、dropout 与版本细节纳入计算图。
主题讲解 · 03:20
FlashAttention 反向传播为何仍是线性额外存储
学习目标
- 能区分前向注意力矩阵与反向注意力梯度矩阵。
- 能解释传统训练为何可能持久化平方级中间量。
- 能说明 FlashAttention backward 如何用分块重计算换存储。
- 能从 dQ、dK、dV 的 shape 判断主要持久梯度规模。
- 能准确限定“关于 N 线性”的假设与统计范围。
- 能说明不同 FlashAttention 版本和框架实现不必采用完全相同的缓存策略。
前置与衔接
前向阶段的主要问题是完整注意力分数或概率矩阵:
反向阶段又会出现同 shape 的上游中间梯度:
因此,只省掉前向 P 还不够。
如果 backward 又把 dS 或 dP 物化回 HBM,平方存储会回来。
视频在 00:30 先给结论:FlashAttention 的前向与反向都能避免平方级持久中间量。
板书先给结论:前向不把完整注意力矩阵写回 HBM,反向也不把其 梯度矩阵写回 HBM。
原视频 · 00:20 ↗核心讲解
1. “线性”首先是关于 N 的说法
视频在 00:05 明确把线性限定为序列长度 N。
单个 head 中:
相应梯度:
当 d 视为固定模型维度时,这些张量关于 N 线性。
而
关于 N 平方。
“线性显存”更严格地应写为 ,不是无条件的 。
2. 传统训练路径为什么有平方中间量
简化前向计算图为
朴素自动微分若为 backward 保存前向激活,可能把 S 或 P 写回 HBM。
反向又计算
和
dP、dS 也是 。
传统训练路径会持久化 S/P 等 中间量及相关梯度,因而平方项可能主导 attention 子层的激活存储。
原视频 · 02:40 ↗3. backward 并不需要把 dP、dS 全部存下来
最终需要交给上一层的是:
它们都是 。
dP 与 dS 只是计算这些最终梯度的中间量。
FlashAttention backward 可以对 Q、K、V 分块:
- 重新读取某个 Q、K tile;
- 在片上重算局部分数 S tile;
- 利用前向保存的每行归一化信息重构局部 P tile;
- 计算局部 dP、dS;
- 立即把它们用于累加 dQ、dK、dV;
- 用完后丢弃局部平方 tile。
因此它计算过平方数量的局部元素,却不让完整 梯度成为持久 HBM 张量。
FlashAttention 的关键不只是省掉前向 A,还通过重计算与分块反向避免持久化 。
原视频 · 00:40 ↗4. 用梯度公式看最终持久张量
由
得到
由
得到
由 得 ,由 得 ;两式展示反向矩阵乘法的 shape 对齐。
原视频 · 02:00 ↗这些公式看起来引用完整 dS 或 P。
但矩阵乘法可以按 tile 做,并把局部贡献规约到最终 dQ、dK、dV。
数学公式定义结果,分块顺序决定内存行为。
5. 反向线性存储的代价是重计算
若前向不保存 P,backward 仍需要 P 才能求 Softmax 梯度。
常见策略是保存每行 log-sum-exp 或等价的归一化统计量,然后重算局部
与
这体现典型的时间—空间权衡:
- 少存平方激活;
- backward 多做一部分局部前向重计算;
- 换来显著减少 HBM 读写与容量压力。
“重计算增加 FLOPs”不必然意味着墙钟时间更差,因为 FlashAttention 的目标正是减少高成本 IO。
实际收益依赖硬件、shape、精度与 kernel 版本。
6. 为什么主要持久状态仍是 Nd
概念上,前向主要持久输入/输出为:
反向主要持久梯度为:
此外还可能保存每行统计量:
局部 S、P、dS、dP tile 的大小由块尺寸控制,不随完整 物化。
FlashAttention 反向以分块重计算换取不存完整 S、P、dS、dP;固定 head dimension d 时,主要持久激活与梯度关于 N 为 。
原视频 · 03:00 ↗7. 矩阵链式法则只是数学入口
视频在 00:56 回顾矩阵乘法链式法则。
矩阵乘积反向需把上游梯度接入对应位置,并按 shape 与原乘法次序处理转置;口诀只能作为记忆辅助。
原视频 · 01:20 ↗它说明 dQ、dK、dV 如何从上游梯度得到。
但一个高性能 backward kernel 还要处理:
- query block 与 key/value block 的循环次序;
- 多个 tile 对 dQ 或 dK 的归约;
- causal mask 与变长序列边界;
- dropout 随机状态;
- mixed precision 累加;
- warp/thread block 映射;
- forward 保存哪些辅助统计量。
因此不能从一条矩阵公式直接推出某个 kernel 的真实峰值显存。
8. 版本与统计口径边界
本课结论适用于 FlashAttention 家族的核心思想:不物化完整 attention 矩阵及其梯度。
但以下内容会因版本或框架而异:
- 前向保存 LSE、m/ 还是其他辅助量;
- backward 的并行划分与归约方式;
- 是否启用 dropout、causal、window attention;
- GQA/MQA 对 K/V head 的共享方式;
- workspace、对齐与临时缓冲大小;
- 自动微分框架外围是否保存额外张量。
“线性额外存储”也不等于“整个模型显存线性”。
参数、优化器状态、MLP 激活、残差、通信缓冲等仍在统计范围之外。
跟练与练习
编者练习
判断哪些张量必须完整持久化 单头 self-attention 中,给出 若 backward 对局部 S、P、dP、dS 即算即用,列出:
- 最终必须输出的梯度;
- 可只以 tile 临时存在的平方中间量;
- 关于 N 的主要持久梯度复杂度。
查看参考答案
最终必须输出:
可按 tile 临时存在:
它们在数学上覆盖 关系,但无需完整写回 HBM。
固定 d 时,最终梯度与每行辅助量的主要持久规模为
这不包含模型其他层、参数、优化器与框架 workspace。
思考题
为什么 backward 重算 P 不会改变数学梯度?
因为在相同 Q、K、mask、缩放与数值规则下,重算得到的是同一前向中间量;它改变保存策略,不改变计算图定义。
若 dropout 存在,则必须能重现相同 dropout mask,不能随意重新采样。
常见误区
误区 1:前向不存 P,反向就无法求梯度
可以利用保存的归一化统计量按块重算 P。
误区 2:计算过 个梯度元素就必须存
错。
局部 tile 可即算即用并立刻规约。
误区 3:反向线性显存意味着反向线性计算
错。
稠密 attention 的 backward 算术量仍含平方级 token 配对。
误区 4:所有 FlashAttention 版本峰值显存完全相同
错。
渐近主项一致不代表常数、workspace 与保存字段一致。
误区 5: 就是整卡训练显存
本课只讨论 attention 算子的主要额外激活/梯度存储。
整模型还包含大量其他状态。
本课小结
- 反向阶段的 dS、dP 也具有 shape,是潜在平方存储源。
- FlashAttention backward 按块重算 S、P、dP、dS,并立即累加到 dQ、dK、dV。
- 完整平方中间梯度无需持久写回 HBM。
- 最终 dQ、dK、dV 均为 ,固定 d 时关于 N 线性。
- 线性额外存储通过重计算和 IO 优化换取,不代表线性算术复杂度。
- 具体保存量与峰值显存仍受 FlashAttention 版本、框架、mask、dropout 与硬件影响。
主题讲解 · 02:48
在线 Softmax 的循环与规约
学习目标
- 能准确区分局部最大值、全局最大值与运行最大值。
- 能定义与最大值绑定的安全 Softmax 分母。
- 能推导两个局部状态的分母重标定公式。
- 能证明顺序循环与分块规约得到同一数学结果。
- 能写出在线 Softmax 状态合并的单位元与结合性直觉。
- 能区分数学等价与有限精度浮点下的逐位一致。
前置与衔接
本地 Whisper 把“局部、最大值、线程、遍历”等识别成近音字,结尾还有超出有效内容的重复短句。 本课术语与公式均按板书及上下文共同校正,不采用字幕尾部重复文本。
视频的核心不是背一个实现细节,而是理解一个可合并状态:
其中:
- m 是当前已见 logits 的最大值;
- 是以 m 为指数基准的安全 Softmax 分母。
视频在 00:05 用“局部信息维护全局信息”概括这件事。
核心讲解
1. 循环与规约处理的是同一种状态
板书在 00:20 区分两种执行方式:
- 单个线程内逐元素循环;
- 线程或块之间做规约。
板书把在线 Softmax 解释为:单个线程内用循环逐步更新状态,线程或块之间再用同一合并规则做规约。
原视频 · 00:20 ↗二者并不是两套不同的 Softmax 公式。
区别只在于数据如何分组:
- 循环:旧状态与一个新元素合并;
- 规约:两个已经聚合好的局部状态合并。
只要合并算子正确,两条路径就恢复同一个全局最大值和全局分母。
2. 每个局部块保存什么
对一个索引集合 A,定义
以及
不是原始指数和
它已经减去局部最大值,从而避免大正数指数溢出。
因此 m 与 必须成对解释。
脱离 m 单独看 ,它没有统一的指数基准。
3. 两个局部最大值如何得到全局最大值
给两个不相交块 A、B:
全局最大值为
这一步很直接。
字幕中的“局部最值再取最值”应准确理解为:
4. 分母为什么不能直接相加
局部分母分别是
二者的基准分别是 与 。
若 ,直接计算
相当于把不同单位的量混在一起。
需要先改写到全局最大值 m 的基准:
因此
合并局部状态 时,先取全局最大值 m,再将每个局部分母乘 后求和。
原视频 · 00:40 ↗所有指数缩放因子都不大于 1:
这保留了安全 Softmax 的数值稳定性。
5. 整体计算例子:[1,2,3,5]
视频在 00:54 取 logits
全局最大值是 5。
安全指数为
因此全局分母是
对 logits ,安全 Softmax 统一减去全局最大值 5,分母为 。
原视频 · 01:00 ↗整行计算给出目标答案。
接下来要证明循环与规约都能恢复同一个 。
6. 顺序循环:先看 [1,2,3],再引入 5
对旧集合 A=[1,2,3]:
新元素单独构成 B=[5]:
单元素块 的局部状态是 、;与旧状态合并时,旧分母必须乘 。
原视频 · 02:00 ↗新的全局最大值:
分母:
展开:
与整行计算完全相同。
7. 并行规约:合并 [1,2] 与 [3,5]
左块 A=[1,2]:
右块 B=[3,5]:
全局最大值仍是
分母合并:
展开:
顺序循环把 与新元素 5 合并;并行规约把 与 合并,两条路径用同一状态算子得到相同结果。
原视频 · 01:40 ↗8. 把合并写成一个算子
定义状态
定义合并
其中
任意分块只要各自保存局部最大值与对应安全分母,就能用相同的重标定公式合并;这正是并行规约成立的基础。
原视频 · 02:20 ↗空集合可以用单位元表示:
因为对任意有效状态 z,数学上有
9. 为什么这个算子满足结合性
每个状态 都等价表示集合 A 的两个真实量:
以及
合并只是把两边的指数和换到共同安全基准后相加。
集合并集本身满足结合律,因此精确实数数学下:
这使 tree reduction、warp reduction 或分层 block reduction 成为可能。
10. 浮点下“相等”不一定逐位相同
实数数学下,循环与任意规约树得到完全相同的 m 与 。 但浮点加法不满足严格结合律, 与 可能在最后几个 bit 上不同。
累加精度、块大小、归约树和编译器优化都会影响误差。
11. 与 FlashAttention 的衔接
在线 Softmax 只需维护每行的 m 与 。
FlashAttention 还要同时合并未归一化的输出分子:
其重标定因子与分母完全一致:
最终
所以本课的 归约,是理解 FlashAttention 输出在线累加的直接前置。
跟练与练习
编者练习
用两种分组验证同一个分母 给定 logits 分别按以下顺序合并: 证明最终状态相同。
- 先合并 [0,2],再合并 [4];
- 先合并 [0],再合并 [2,4]。
查看参考答案
第一种分组:
合并得
第二种分组:
合并得
两种分组都等于整行以最大值 4 为基准的安全分母。
常见误区
误区 1:局部分母可以直接相加
只有局部最大值相同,或都已经统一到同一全局最大值时才可以。
一般情况必须乘 。
误区 2:运行最大值改变时,只更新新元素
错。
旧分母的所有贡献都要一起缩放到新基准。
误区 3:“动态规划”意味着要保存长度为 N 的 DP 表
本课只需常数个每行状态。
“动态规划”是局部状态递推的理解方式。
误区 4:规约与循环必须逐位相同
数学结果相同,浮点执行顺序不同可能产生微小舍入差异。
误区 5:最大值 m 本身就是 Softmax 分母
m 只是数值稳定性的指数基准。
真正分母是与 m 绑定的 。
误区 6:板书中的 sum 是普通 logits 求和
不是。
它表示安全指数之和:
本课小结
- 在线 Softmax 的每行状态是运行最大值 m 与安全分母 。
- 两块合并时先取 。
- 分母按 重标定。
- 顺序循环是旧状态与新元素合并,规约是两个局部状态合并。
- 精确数学下合并算子满足结合律,因此可并行规约。
- 浮点下不同规约树可能不逐位相同,但应实现同一数学 Softmax。
- 这一 状态合并是理解 FlashAttention 输出在线累加的基础。
单元综合
从分块等价到在线归约:FlashAttention 的前向与反向
单元能力目标
完成本单元后,应能把 FlashAttention 还原为一组可验证的数学等价变换与存储调度。
具体需要做到:
- 证明按 行和 行分块不会改变注意力矩阵的定义;
- 从安全 Softmax 推导在线状态 ;
- 用可合并的归约解释顺序扫描、并行归约与分块计算;
- 说明 FlashAttention 如何在片上存储分数块,并立即并入输出累加器;
- 判断 causal mask 中哪些 tile 能整块跳过;
- 区分 FlashAttention 版本调度变化与核心数学不变量;
- 区分线性持久存储、平方级 dense 计算与实际 HBM 流量;
- 用矩阵链式法则写出 ;
- 解释反向为何能通过重计算避免持久化 中间梯度。
概念连接
1. 从完整注意力定义出发
对单个 attention head,先写出
若有 causal mask ,则实际分数为
其中未来位置对应 。
FlashAttention 没有改变这三个数学对象的定义,改变的是它们被计算、保存和丢弃的顺序。
2. 每个分数元素只依赖一对行向量
分数矩阵的元素是
它只依赖 的第 行和 的第 行。
所以,只要所有 组合都被计算且不重不漏,就能恢复与整体矩阵乘法相同的 。
3. 按 Q 行分块是纵向拼接
若
则矩阵乘法的分块法则直接给出
每个 只负责一组 query 行,但要遍历它们可见的全部 key,才能得到这些行的全局 Softmax 统计量。
4. 同时分 Q 块与 K 块
再将 分为 ,某个 tile 是
所有 tile 构成原分数矩阵的完整笛卡尔覆盖。
但此时一个 对应的 Softmax 分母被分散在多个 中,需要在遍历 key tile 时逐块合并。
5. 安全 Softmax 避免指数溢出
对一行分数 ,直接计算 可能溢出。
安全写法为
传统的逻辑数组遍历是:
- 第一遍找 ;
- 第二遍计算 ;
- 第三遍输出归一化概率。
这里的“三遍”是逻辑遍历,并不等同于某个具体 kernel 必然产生三次完整 HBM 读写。
6. 在线 Softmax 同时更新最大值和分母
扫描到新元素 时,令
原来的分母是围绕旧最大值 表示的,需换到新基准 :
因此 就是一行 Softmax 的最小核心状态。
第一次扫描可以同时得到安全的 和 ,再扫描一次归一化,将三遍降为两遍。
7. 两段状态可以稳定合并
设两段分数的状态分别是 与 。
合并时:
这正是将两段的安全指数和改写到同一个全局最大值下。
顺序在线循环可看作“已有状态”与“单个元素状态”反复合并;分块并行则是“块状态”之间的树形归约。
8. 数学可结合不代表浮点逐位一致
在精确实数算术中,上述合并对分段方式是结合的。
在浮点实现中,不同归约树会改变加法顺序,可能产生微小舍入差异。
所以需区分:
- 数学上计算同一个 Softmax;
- 数值上不承诺不同 kernel 逐 bit 一致。
9. FlashAttention 还需第三类状态
最终输出的一行是
定义未归一化分子累加器
当最大值从 变为 时, 和 都必须乘同一重标定因子
于是 FlashAttention 每行维护三类概念状态:运行最大值、安全分母与输出分子。
10. 前向循环的不变量
固定一个 query tile 后,遍历所有可见的 :
- 把 所需部分携入片上存储;
- 计算局部分数 ;
- 应用 scale 和 mask;
- 计算局部最大值与指数和;
- 重标定旧状态,并入新块;
- 用局部权重与 更新输出累加器;
- 丢弃局部 ,进入下一个 key tile。
遍历完成后再用 归一化。
关键不变量是:处理完前 个 key tile 后,状态已精确表示这些 tile 联集上的 Softmax 与输出分子。
11. 局部分数矩阵不落入 HBM
普通实现若把整个 或 写回 HBM,会产生 级别的持久存储和读写。
FlashAttention 让局部 只在 SRAM/寄存器等片上层级短暂存活,它与 的贡献并入输出状态后即可丢弃。
这是 IO-aware 调度:优化目标不是少算所有 dense 内积,而是避免反复物化并搬运巨大中间矩阵。
12. Causal mask 可以消除完全未来 tile
在 decoder causal attention 中,对于 query 位置 ,key 位置 必须被屏蔽。
对一个 tile:
- 若其所有 都大于所有对应 ,整块都是未来,可直接跳过;
- 若 tile 与对角线相交,内部同时存在可见与不可见元素,仍需细粒度 mask;
- 若 tile 完全在对角线下方,可整块正常计算。
跳过完全未来 tile 是精确消除零贡献,不是近似。
13. 单 token Decode 不等于普通 Prefill 的上三角
自回归 decode 中,新 query 位于当前最后,所有已缓存 key 都是它的过去或当前位置。
此时不存在同样的大面积未来区域。
多 token prefill 或块状验证则需在新块内保留 causal 约束。
因此“能跳过上三角”需结合 query/key 范围和实际 tile 边界判断。
14. FlashAttention-2 的外 Q 内 K/V 调度
观察到的高层循环是:
- 外层固定 与其输出 ;
- 内层流式遍历 ;
- 尽量留在片上;
- 内层完成后一次写回 。
与高层上外 内 的安排相比,这能减少对 query、输出与在线状态的重复读写。
但它不表示 只从 HBM 读一次,也不表示循环交换是 v2 的唯一优化。
实际 tile 大小、warp 分工、共享存储与并行策略还受 GPU 架构和 kernel 版本影响。
15. 平方分数矩阵是持久存储的压力源
对
分数和概率矩阵都是
若持久化 或 ,元素数为 。
而 各自为 元素。
对比一个 张量,比值为
例如 时,比值约为 。
16. 线性显存不意味线性计算
FlashAttention 的关键持久张量可随 增长,并不需保存整个 分数矩阵。
但 dense attention 仍需计算大量 query-key 内积和权重-value 累加,主要算术量仍是
所以需同时报告:
- 持久中间存储;
- HBM 读写量;
- 片上临时存储;
- 计算量;
- 硬件上的实际吞吐。
17. 反向先沿 O=PV 传播
设损失对输出的梯度为 。
对
矩阵微分是
用 Frobenius 内积或 trace 把微分中的待求项移到末尾,得
这里的乘法次序不能交换。
18. Softmax 反向是逐行 JVP
对某一行 ,已知上游梯度 ,有
即先求该行的标量
再计算
scale 、mask 与 dropout 都必须按前向的真实路径加入,不能在简化推导后遗漏。
19. 再沿 S=QK^T 传播
对
有
若前向包含 ,相应梯度也需乘该系数。
用 shape 可以自检:
20. 反向不必持久保存平方梯度
朴素理解可能会把 都当作需要常驻的 张量。
FlashAttention 反向逐块重建局部中间量:
- 用 重算局部分数;
- 借助前向保存的行统计量重建局部 ;
- 计算局部 ;
- 立即累加到 ;
- 丢弃局部平方 tile。
最终梯度 与输入一样为 。
线性额外存储是用重计算、分块调度和 IO 交换得到的,并不意味 backward 变成线性算术量。
对比与决策
1. 数学等价与逐位一致
- 分块覆盖同一组 内积,并用正确在线状态合并:数学目标相同。
- 归约顺序和 kernel 不同:浮点结果可能有微小差异。
- 漏 tile、重 tile、忘记重标定或 mask 错位:不再数学等价。
2. 显存复杂度与计算复杂度
- 不持久化 中间量:可将关键额外存储降到线性量级。
- 仍计算 dense query-key 组合:算术仍是平方级 token 耦合。
- 运行更快:主要来自减少 HBM 物化/搬运和提高片上复用,不能仅由大 O 记号推断。
3. 哪些 causal tile 能跳过
- 完全位于未来区:整块跳过。
- 跨越对角线:计算 tile,内部细粒度 mask。
- 完全位于可见区:正常计算。
- 单 token decode:已有 key 通常全可见,不套用普通 prefill 的上三角图像。
4. 前向存储还是反向重计算
- 保存更多 :反向重算少,但显存和 IO 大。
- 保存紧凑行统计:反向逐块重构,显存少,增加重算与调度复杂度。
选择应根据显存瓶颈、算力余量、硬件带宽与 kernel 可用性,不是只比较某一个公式。
5. v1/v2 调度与不变的数学核心
高层循环顺序、并行分解和 warp 工作分配可以改变,但以下原则不变:
- 完整覆盖所有可见分数;
- 不同块的 Softmax 统计使用全局最大值重标定;
- 输出分子与分母同步合并;
- 局部平方 tile 不作为全局持久张量。
综合训练
编者练习 1
设 ,。写出 的 分块形式,并说明为什么单独计算 后就立即对其行做 Softmax 一般是错的。
查看参考答案
分块矩阵是
的每一行需要与 和 中所有可见 key 共同归一化。只对第一个 tile 单独 Softmax 会使其分母缺少 的指数项。正确做法是在两个 key tile 之间合并 。
编者练习 2
两个分数块的 Softmax 状态为 求合并后的 。
查看参考答案
全局最大值是 。
第一块的分母需从基准 2 换到基准 4:
数值约为 。若有输出分子状态,第一块的 也需乘 ,第二块乘 。
编者练习 3
对 的单头注意力,与一个 张量相比, 分数矩阵的元素数多少倍?这个数字能否证明 FlashAttention 的计算复杂度是线性?
查看参考答案
比值为
这只比较两类张量的元素数,说明不物化分数矩阵的存储优势。Dense attention 仍需处理约 个 query-key 组合,主要算术量仍为 。
编者练习 4
已知 写出不经 Softmax 这一步时的 ,并用 shape 检查结果。
查看参考答案
设 ,则
将逐行 Softmax JVP 得到的梯度记为 ,则
右侧结果分别与 的 shape 一致。若 含 scale,还需传播 。
进入下一单元前
- 已能从 证明分块覆盖的完整性。
- 已能独立推导在线 Softmax 的 更新与块合并公式。
- 已能说明输出分子累加器为什么必须与分母同步重标定。
- 已能判断 causal tile 的整块跳过与块内 mask 边界。
- 已能区分 持久存储与 dense 计算。
- 已能写出 attention 反向的矩阵梯度链条。
- 已能解释 backward 通过逐块重计算避免持久化平方中间量。
- 若仍把“分块”理解为近似,回看 P26、P27。
- 若仍会在块内独立做 Softmax,回看 P28、P33、P36。
- 若仍把线性显存写成线性计算,回看 P31、P32、P35。
- 若对反向乘法次序不确定,回看 P34,并每次做 shape 自检。