LLM WIKI · 课程精读

LEARNING UNIT · 12

在线 Softmax 与 FlashAttention

用分块等价性和在线归约状态解释 FlashAttention 的前向、反向、掩码与显存复杂度。

已整理章节
12 节
单元来源
11 条视频
总时长
32:04
状态
已发布
学习位置
12 / 20
01

主题讲解 · 02:32

为什么分块计算仍得到同一个注意力矩阵

学习目标

  • 能从单个元素 AijA_{ij} 解释 QKTQK^T 的含义。
  • 能写出 Q 与 K 分块后每个输出块的公式。
  • 能证明分块只改变计算顺序,不改变注意力分数矩阵。
  • 能说明 HBM、SRAM 与线程块在分块计算中的角色。
  • 能指出“矩阵乘法等价”与“完整 FlashAttention 正确性”之间还差什么。

前置与衔接

需要会做矩阵乘法,并知道朴素 attention 的分数矩阵来自 QKTQK^T

本课是“在线 Softmax 与 FlashAttention”单元的入口。先证明 tile 计算不会改写 QKTQK^T,再学习在线 softmax 如何在不保存完整分数矩阵时得到同一输出。

核心讲解

1. 朴素算法先物化完整分数矩阵

忽略缩放因子 1/dk1/\sqrt{d_k},注意力分数为

A=QKT.A=QK^T.
图 1

朴素注意力先计算完整的 A=QKᵀ,示例矩阵给出可逐项核对的结果。

原视频 · 00:20 ↗

Q,KRN×dQ,K\in\mathbb{R}^{N\times d},则 ARN×NA\in\mathbb{R}^{N\times N}

朴素实现通常把这个完整 N×NN\times N 中间量写到高带宽显存 HBM,再读取它做 softmax 和后续乘 VV

当序列长度 NN 增大时,这个中间量的元素数按 N2N^2 增长。

FlashAttention 的核心工程目标不是修改 attention 定义,而是减少这个中间量在 HBM 中的读写和物化。

2. 每个输出元素本来就是一个局部内积

矩阵乘法定义给出

Aij=Qi,:Kj,:T.A_{ij}=Q_{i,:}K_{j,:}^T.
图 2

注意力分数 Aᵢⱼ 等于 Q 的第 i 行与 K 的第 j 行做内积。

原视频 · 00:40 ↗

这里 Qi,:Q_{i,:} 是 Q 的第 ii 行,Kj,:TK_{j,:}^T 是 K 的第 jj 行转置成的列向量。

展开后

Aij=r=1dQirKjr.A_{ij}=\sum_{r=1}^{d}Q_{ir}K_{jr}.

这条公式非常重要:任何分块方案只要让每个 (i,j)(i,j) 最终累加完全相同的 dd 个乘积,结果就不会改变。

所以正确性检查不需要迷信“FlashAttention”这个名字,只要回到每个输出元素的求和项即可。

3. Q 按行切,Kᵀ 按列切

把 Q 按 token 行切成两个块:

Q=[Q1Q2].Q=\begin{bmatrix}Q_1\\Q_2\end{bmatrix}.

K 原本也按 token 行切:

K=[K1K2],K=\begin{bmatrix}K_1\\K_2\end{bmatrix},

转置后就成为列块:

KT=[K1TK2T].K^T=\begin{bmatrix}K_1^T & K_2^T\end{bmatrix}.
图 3

Q 沿行方向切块,Kᵀ 沿列方向切块,二者组合覆盖完整的注意力矩阵。

原视频 · 01:00 ↗

因此

QKT=[Q1K1TQ1K2TQ2K1TQ2K2T].QK^T= \begin{bmatrix} Q_1K_1^T & Q_1K_2^T\\ Q_2K_1^T & Q_2K_2^T \end{bmatrix}.

每个 QiKjTQ_iK_j^T 都是完整注意力矩阵中的一个矩形 tile。

4. 拼回所有 tile 就是原矩阵

课程用具体数字逐块计算,发现每个 tile 与朴素结果中相同位置的小块完全一致。

图 4

每个输出块 QᵢKⱼᵀ 仍由相同的行列内积组成,拼接后与朴素结果一致。

原视频 · 01:20 ↗

原因并不神秘。

设第 ii 个 Q 块覆盖行集合 IiI_i,第 jj 个 K 块覆盖行集合 JjJ_j

块乘法 QiKjTQ_iK_j^T 中任意元素仍是

r=1dQurKvr,uIi, vJj.\sum_{r=1}^{d}Q_{ur}K_{vr},\qquad u\in I_i,\ v\in J_j.

这与朴素矩阵中 AuvA_{uv} 的定义一字不差。

分块改变的是:

  • 先算哪些行列区域;
  • 一个 tile 何时搬进片上存储;
  • 中间结果是否立即写回 HBM。

分块没有改变的是:

  • 每个输出元素对应的 Q 行与 K 行;
  • 内积的收缩维度 dd
  • 每个内积包含的乘加项。

因此,在相同浮点精度与舍入顺序假设下,分块结果与朴素结果相同;实际浮点实现可能有极小舍入差异,但算法语义等价。

5. GPU 上为什么值得分块

HBM 容量大但访问代价高,片上 SRAM 容量小但带宽高。

FlashAttention 让线程块负责一个或若干 tile:

  1. 把 Q tile 从 HBM 搬到 SRAM;
  2. 分批把 K/V tile 搬到 SRAM;
  3. 计算局部 QiKjTQ_iK_j^T
  4. 立即把局部结果并入 softmax 状态和输出累加;
  5. 不把完整 N×NN\times N 分数矩阵长期写回 HBM。
图 5

GPU 线程块把 Q/K tile 从 HBM 搬到 SRAM 计算局部注意力块,而 softmax 仍需要跨块的全局行信息。

原视频 · 02:20 ↗

课程在结尾指出一个关键困难:softmax 是逐行的,但每个线程块一次只看到一行的一部分。

这意味着证明 QKTQK^T 分块等价还不够,还要证明局部 softmax 状态可以正确合并。

6. 编者补充:完整正确性还需要在线 softmax

对一行分数 ss,稳定 softmax 通常使用全局最大值

m=maxjsjm=\max_j s_j

和归一化分母

=jesjm.\ell=\sum_j e^{s_j-m}.

若当前只看到一个分块 s(b)s^{(b)},可先得到局部最大值 mbm_b 与局部分母 b\ell_b

把旧状态与新块合并时,先更新

mnew=max(mold,mb),m_{\text{new}}=\max(m_{\text{old}},m_b),

再把二者缩放到同一基准:

new=emoldmnewold+embmnewb.\ell_{\text{new}} =e^{m_{\text{old}}-m_{\text{new}}}\ell_{\text{old}} +e^{m_b-m_{\text{new}}}\ell_b.

输出的加权和也做同样重标定,最终再除以 new\ell_{\text{new}}

因此完整 FlashAttention 正确性包含两层:

  1. 分块矩阵乘法覆盖同一组注意力分数;
  2. 在线 softmax 用可合并状态恢复与全局 softmax 相同的归一化输出。

本条视频完整说明了第一层,并在结尾引出第二层;它不是对整个 FlashAttention 前向算法的全部证明。

跟练与练习

原视频练习

编者练习

Q 有 8 行、K 有 12 行,二者特征维度都是 64。若 Q 每 4 行一块,K 每 3 行一块,最终会得到多少个注意力 tile?每个 tile 的 shape 是什么?

查看参考答案

Q 有 8/4=28/4=2 个行块,K 有 12/3=412/3=4 个行块,因此输出有 2×4=82\times4=8 个 tile。每个 tile 是 QiKjTQ_iK_j^T,shape 为 4×34\times3;拼接后得到 8×128\times12 的完整分数矩阵。

编者练习 2

为什么“每个输出 tile 都正确”仍不足以证明 FlashAttention 的最终输出正确?

查看参考答案

attention 还要对整行分数做 softmax,再乘 V。单个 tile 只包含一行的局部分数,缺少全局最大值和全局归一化分母。还必须证明在线 softmax 的最大值、分母和加权输出状态可以跨 tile 正确合并。

常见误区

  • 误区:K 横着切,所以 Kᵀ 仍横着切。纠正:转置后 token 行块会变成列块。
  • 误区:分块近似了 QKTQK^T。纠正:它重排精确乘加,算法语义并非低秩或稀疏近似。
  • 误区:FlashAttention 把注意力复杂度从 N2N^2 变成线性。纠正:计算量通常仍是二次级,主要降低的是中间显存与 HBM I/O。
  • 误区:证明块矩阵乘法就完成了全部证明。纠正:还需要在线 softmax 与乘 V 的合并正确性。
  • 误区:浮点结果必须逐 bit 相同。纠正:重排归约顺序可能产生微小舍入差,但应在数值容差内等价。

本课小结

  • AijA_{ij} 始终是 Q 的一行与 K 的一行做内积。
  • Q 行块与 Kᵀ 列块的笛卡尔组合覆盖完整注意力矩阵。
  • tile 计算改变执行顺序与内存流量,不改变每个输出元素的定义。
  • GPU 分块的收益来自复用片上 SRAM、减少完整分数矩阵的 HBM 读写。
  • 下一步要学习在线 softmax 如何合并局部最大值、归一化分母与加权输出。
02

主题讲解 · 02:31

按 Q 行分块为什么仍等于完整注意力分数

学习目标

  • 能从元素定义解释注意力分数矩阵。
  • 能证明只沿 Q 的行轴分块不改变 QKQK^\top
  • 能区分“结果完全相同”与“执行代价一定更低”。
  • 能说明完整 key 轴为何提供一行 Softmax 的全局统计量。
  • 能识别视频中 CPU/缓存示意与具体实现之间的边界。

前置与衔接

注意力在忽略缩放、掩码和多头维度后,可先写成

A=QK.A=QK^\top.

视频用一个称为 SlimAttention 的方案说明:把 QQ 沿 token 行切成 tile,但在当前计算层次让 KK^\top 保持覆盖完整 key 轴。

问题是,分开计算每个 query tile,最后还能否严格得到原来的 AA

本课先证明线性代数等价,再讨论这一分块对 Softmax 与内存层次意味着什么。

核心讲解

1. 元素级定义不随分块改变

QRNq×d,KRNk×d.Q\in\mathbb{R}^{N_q\times d},\qquad K\in\mathbb{R}^{N_k\times d}.

A=QKRNq×Nk.A=QK^\top\in\mathbb{R}^{N_q\times N_k}.

i,ji,j 个元素为

Aij=qikj.A_{ij}=q_i^\top k_j.
图 1

注意力分数矩阵的元素 AijA_{ij} 等于第 ii 个 query 行向量与第 jj 个 key 行向量的内积。

原视频 · 00:20 ↗

图中选择一行 qiq_iKK^\top 的一列,也就是原 KK 的一行 kjk_j

二者内积填回 AijA_{ij}

只要分块后仍计算同一对 qi,kjq_i,k_j,该元素的值就不会改变。

2. 沿 Q 的行轴切成两个 tile

QQ 按行分成

Q=[Q1Q2],Q= \begin{bmatrix} Q_1\\ Q_2 \end{bmatrix},

其中 Q1Q_1Q2Q_2 分别包含若干完整 query 行。

图 2

视频中的 SlimAttention 方案沿 token 行把 QQ 切成多个 tile,而 KK^\top 在该层解释中保持整块。

原视频 · 00:40 ↗

这里没有沿特征维 dd 切断单个 token 向量。

每个 query tile 仍包含完整的 dd 维向量,因此可以与每个 key 做完整内积。

3. 分块乘法的直接证明

矩阵乘法对纵向拼接满足

[Q1Q2]K=[Q1KQ2K].\begin{bmatrix} Q_1\\ Q_2 \end{bmatrix} K^\top = \begin{bmatrix} Q_1K^\top\\ Q_2K^\top \end{bmatrix}.

左侧是一次计算完整 QKQK^\top

右侧是分别计算两个行块,再按原顺序纵向拼接。

图 3

Q1KQ_1K^\topQ2KQ_2K^\top 分别计算后按行拼接,恰好恢复完整的 QKQK^\top

原视频 · 01:20 ↗

两者逐元素相同,不依赖近似,也不依赖浮点外的特殊假设。

在真实浮点执行中,不同 kernel 的归约顺序可能造成末位舍入差异;这里的“等价”指相同数学表达式,不承诺 bitwise identical。

4. 为什么不能把结果横向拼错

Q1Q_1Q2Q_2 是沿行轴切分,所以各自产物 shape 为

QrKRBr×Nk.Q_rK^\top\in\mathbb{R}^{B_r\times N_k}.

它们都覆盖完整 key 列,只覆盖不同 query 行。

因此必须沿行轴拼接。

若错误地沿列轴拼接,shape 和 AA 的 query/key 语义都会改变。

判断拼接方向时,应始终追踪“被切的是哪一根轴”。

5. 一个具体 shape 例子

假设

QR4×4,KR4×4.Q\in\mathbb{R}^{4\times4},\qquad K^\top\in\mathbb{R}^{4\times4}.

QQ 切成两个 2×42\times4 tile:

Q1,Q2R2×4.Q_1,Q_2\in\mathbb{R}^{2\times4}.

Q1K,Q2KR2×4.Q_1K^\top,Q_2K^\top\in\mathbb{R}^{2\times4}.

按行拼接得到 4×44\times4,与完整乘法 shape 一致。

视频板书中的数字例子还逐项验证了对应元素相同。

6. 从数学分块到执行分块

视频随后给出内存示意:Q、K 位于较慢的主存层,执行单元把当前 Q tile 与所需 K 数据搬入更快的片上存储,再计算一个 A 的行块。

图 4

一个执行单元把当前 query tile 与所需 key 数据从较慢内存搬到片上快速存储,再计算对应分数行块。

原视频 · 01:40 ↗

一个线程或线程块负责哪个 tile、K 是否一次完整驻留,都取决于硬件与实现。

数学证明只要求每个所需内积最终被计算,不能从证明本身推出某种固定调度一定最快。

7. 为什么完整 key 轴对 Softmax 有利

注意力 Softmax 沿 key 轴逐行计算:

Pij=exp(Aijmi)j=1Nkexp(Aijmi),mi=maxjAij.P_{ij} = \frac{\exp(A_{ij}-m_i)} {\sum_{j'=1}^{N_k}\exp(A_{ij'}-m_i)}, \qquad m_i=\max_j A_{ij}.

若一个 Q tile 与完整 KK^\top 相乘,所得行块覆盖每一行的全部 NkN_k 个 key 分数。

图 5

当前 query tile 覆盖完整 key 轴时,得到的是完整分数行块,可直接获得该行 Softmax 的最大值与分母。

原视频 · 02:20 ↗

于是该 tile 内可以直接得到每行的全局最大值 mim_i 与完整分母。

不必跨多个 K tile 逐次修正同一行的 Softmax 状态。

8. 与 FlashAttention 双轴分块的差别

FlashAttention 通常同时对 query 轴与 key 轴分块。

单个分数 tile 只看到部分 key,因此没有一整行的全局最大值与分母。

它需要维护在线 Softmax 状态,在新 K/V tile 到来时重标定旧的部分结果。

视频中的 SlimAttention 行分块则让一个 Q tile 在该层次覆盖完整 K 轴,所以行归一化更直接。

但完整 K 是否能放进快速存储、是否需要在更低层继续切块,是实现层问题。

9. 正确性不等于任何尺寸都高效

沿 Q 行分块永远保持矩阵乘法的数学正确性。

性能却取决于:

  • Nk×dN_k\times d 的 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 与 PVPV
  • 面向 CPU 还是 GPU;
  • 支持哪些 dtype 与序列长度。

跟练与练习

原视频练习

编者练习

QR6×8Q\in\mathbb{R}^{6\times8}KR10×8K\in\mathbb{R}^{10\times8},把 Q 沿行切成大小 2、4 的两个 tile。两个局部乘积的 shape 分别是什么,怎样恢复完整 A?

查看参考答案

KR8×10K^\top\in\mathbb{R}^{8\times10}。两个 Q tile 的 shape 分别为 2×82\times84×84\times8,乘积分别为 2×102\times104×104\times10。沿行轴纵向拼接后得到 6×106\times10,即完整 QKQK^\top

编者练习 2

为什么沿 key 轴把同一分数行切成两块后,不能各自独立做 Softmax 再直接拼接?

查看参考答案

Softmax 的分母和最大值必须覆盖该 query 行的全部 key。两块各自归一化会让每块权重和都等于 1,拼接后不再是完整行的 Softmax。必须先合并全局统计量,或使用在线 Softmax 的重标定公式。

常见误区

  • 误区:分块会近似矩阵乘法。纠正:按行分块并完整计算所有内积时,数学结果严格相同。
  • 误区:Q 按行切后应横向拼接结果。纠正:切的是 query 行轴,结果也沿行轴纵向拼接。
  • 误区:一个 tile 可以切断 token 的特征维却仍直接套同一证明。纠正:沿收缩维切分时还需要对部分内积求和。
  • 误区:完整 K 意味着所有实现都把整个 K 一次放进 SRAM。纠正:视频描述的是当前抽象层,实际可能有下一级分块。
  • 误区:有完整分数行就完全没有缓存成本。纠正:仍需考虑 K 搬运、Q tile、输出和硬件容量。
  • 误区:数学等价必然带来性能提升。纠正:性能还取决于数据复用、带宽与并行度。

本课小结

  • Aij=qikjA_{ij}=q_i^\top k_j 给出分块正确性的元素级依据。
  • Q=[Q1;Q2]Q=[Q_1;Q_2] 按行切分,有 QK=[Q1K;Q2K]QK^\top=[Q_1K^\top;Q_2K^\top]
  • 当前 Q tile 覆盖完整 key 轴时,可直接获得每行 Softmax 的全局统计量。
  • 分块把数学计算映射到分层存储,但正确性与性能是两个结论。
  • 视频的 SlimAttention、CPU 与 SRAM 表述属于具体示意,实际实现需按论文、代码和硬件复核。
03

主题讲解 · 03:44

在线 Softmax 如何把安全计算从三遍降到两遍

学习目标

  • 能区分朴素 Softmax、安全 Softmax 与在线 Softmax。
  • 能解释减去最大值为何不改变概率且提高数值稳定性。
  • 能推导在线更新中的最大值与分母重标定公式。
  • 能说明安全实现的前三个逻辑阶段如何融合为两遍。
  • 能区分算法遍历数、kernel 数与真实内存流量。

前置与衔接

对向量 x=(x1,,xN)x=(x_1,\ldots,x_N),Softmax 定义为

pi=exij=1Nexj.p_i=\frac{e^{x_i}}{\sum_{j=1}^{N}e^{x_j}}.

直接实现需要先得到分母,才能写出每个 pip_i

安全 Softmax 还要先求全局最大值,避免指数溢出,于是看起来需要三次遍历。

在线 Softmax 的关键是:最大值变化时,把此前累计的分母转换到新的指数基准,从而在同一遍流式扫描中同时维护最大值与分母。

核心讲解

1. 朴素 Softmax 的两阶段依赖

以视频中的

x=[1,2,3,5]x=[1,2,3,5]

为例。

图 1

朴素 Softmax 先计算指数和,再用同一分母归一化各元素,因此至少包含归约与归一化两个阶段。

原视频 · 00:20 ↗

第一遍计算

S=jexj.S=\sum_j e^{x_j}.

第二遍再写出

pi=exi/S.p_i=e^{x_i}/S.

在不知道 SS 之前,不能得到最终归一化结果。

因此朴素的独立数组实现至少有“归约 + 写出”两个逻辑阶段。

2. 为什么直接取指数可能不安全

若某个 xix_i 很大,exie^{x_i} 可能超出浮点数可表示范围。

即使最终比例本应有限,中间指数也可能先变成正无穷,导致不定比值或 NaN。

Softmax 对整体平移不变:

exicjexjc=exiececjexj=exijexj.\frac{e^{x_i-c}}{\sum_j e^{x_j-c}} = \frac{e^{x_i}e^{-c}}{e^{-c}\sum_j e^{x_j}} = \frac{e^{x_i}}{\sum_j e^{x_j}}.

所以可以选择 c=maxjxjc=\max_j x_j

3. 安全 Softmax 把最大指数压到 1

视频示例的全局最大值为 5。

减去它后得到

[1,2,3,5]5=[4,3,2,0].[1,2,3,5]-5=[-4,-3,-2,0].
图 2

安全 Softmax 将 [1,2,3,5] 减去全局最大值 5,得到 [-4,-3,-2,0],结果不变但指数不再溢出。

原视频 · 01:20 ↗

此时最大的指数是 e0=1e^0=1,其他指数都位于 (0,1](0,1],显著降低上溢风险。

画面也证实了 ASR 中漏掉的负号:移位向量是“[-4,-3,-2,0]”,不是“[4,-3,-2,0]”。

4. 朴素安全实现为何是三遍

安全 Softmax 可按以下方式实现:

  1. 第一遍:求 m=maxixim=\max_i x_i
  2. 第二遍:求 =iexim\ell=\sum_i e^{x_i-m}
  3. 第三遍:写出 pi=exim/p_i=e^{x_i-m}/\ell
图 3

朴素安全实现依次求全局最大值、求移位指数和、写出归一化结果,共三次逻辑遍历。

原视频 · 02:00 ↗

第二遍依赖第一遍的 mm,第三遍又依赖第二遍的 \ell

若把三步机械实现为三次读取输入,就会产生三遍数据访问。

5. 在线状态只需维护两个量

在线 Softmax 处理到前 tt 个元素时维护

mt=max1itxim_t=\max_{1\le i\le t}x_i

t=i=1teximt.\ell_t=\sum_{i=1}^{t}e^{x_i-m_t}.
图 4

在线 Softmax 为每个分块维护局部最大值 mm 与相对该最大值的指数和 \ell

原视频 · 02:40 ↗

mtm_t 是当前前缀最大值。

t\ell_t 不是原始指数和,而是所有已见元素相对于当前最大值 mtm_t 的指数和。

这两个量足以表示已经扫描部分的安全归一化信息。

6. 新元素到来时怎样更新

读入 xtx_t 后,先更新

mt=max(mt1,xt).m_t=\max(m_{t-1},x_t).

旧分母 t1\ell_{t-1} 是以 mt1m_{t-1} 为基准。

要换成新基准 mtm_t,必须乘 emt1mte^{m_{t-1}-m_t},于是

t=t1emt1mt+extmt.\ell_t = \ell_{t-1}e^{m_{t-1}-m_t} +e^{x_t-m_t}.

若新元素没有刷新最大值,旧分母系数为 1。

若新元素更大,旧的所有指数项会统一缩小到新基准。

7. 两个分块如何合并

设分块 a、b 分别维护

(ma,a),(mb,b).(m_a,\ell_a),\qquad(m_b,\ell_b).

先取

m=max(ma,mb).m=\max(m_a,m_b).

再把两个局部分母都变换到基准 mm

=aemam+bembm.\ell = \ell_a e^{m_a-m} +\ell_b e^{m_b-m}.
图 5

合并分块时先取新最大值,再用 embme^{m_b-m} 重标定各局部分母,保证它们处于同一指数基准。

原视频 · 03:00 ↗

这就是板书中用两个局部向量信息更新全局 max 与 sum 的公式。

它允许树形归约、线程块局部归约或流式扫描,不要求先单独遍历完整向量求最大值。

8. 用 [1,2] 和 [3,5] 验证

左块:

ma=2,a=e12+e22=e1+1.m_a=2,\qquad \ell_a=e^{1-2}+e^{2-2}=e^{-1}+1.

右块:

mb=5,b=e35+e55=e2+1.m_b=5,\qquad \ell_b=e^{3-5}+e^{5-5}=e^{-2}+1.

合并最大值为 m=5m=5,所以

=(e1+1)e25+(e2+1)e55=e4+e3+e2+1.\ell =(e^{-1}+1)e^{2-5} +(e^{-2}+1)e^{5-5} =e^{-4}+e^{-3}+e^{-2}+1.

这正好等于直接对“[-4,-3,-2,0]”求指数和。

9. 为什么从三遍降到两遍

在线更新把“求最大值”和“求相对指数和”融合到第一次扫描:

x(m,).x\longrightarrow(m,\ell).

第二遍再根据最终 m,m,\ell 写出

pi=exim/.p_i=e^{x_i-m}/\ell.
图 6

最大值与分母可在一次流式遍历中共同更新,第二遍只负责归一化输出,从三遍降为两遍。

原视频 · 03:20 ↗

所以独立 Softmax 数组实现从三次逻辑遍历变成两次。

少掉的是单独求最大值的那一遍,不是消除最终归一化。

10. 数值稳定性没有被牺牲

在线状态始终以“当前已见最大值”为指数基准。

新最大值只会不减,所有指数的自变量都满足

ximt0.x_i-m_t\le0.

因此不会为少一次遍历而恢复到直接计算巨大指数的不安全做法。

有限精度下,不同归约树仍会产生微小舍入差异,但算法维持了 max-shift 的稳定原则。

11. 循环数不等于 kernel 数

视频用 for 循环与 HBM 到 SRAM 搬运解释收益。

这是很好的数据流直觉,但工程上还要区分:

  • 源代码循环次数;
  • GPU kernel 启动次数;
  • 一个 kernel 内的线程归约;
  • 编译器是否已经融合;
  • 数据是否真正离开 cache 或寄存器。

若安全 Softmax 已由高度融合 kernel 实现,不能机械地把“三个公式阶段”换算成“三次完整 HBM 往返”。

在线公式提供可融合性,真实性能仍需测量。

12. 与 FlashAttention 的关系

FlashAttention 不把完整分数行存到 HBM 后再跑独立 Softmax。

它在遍历 score tiles 时维护每行的运行最大值、分母和输出累加器。

新 tile 改变最大值时,不仅要重标定 \ell,还要重标定此前的输出部分和。

因此在线 Softmax 是 FlashAttention 分块等价性的核心组件,但实际 kernel 比本课的独立向量两遍算法更进一步融合了 PVPV

13. 论文与版本边界

视频把在线 Softmax 追溯到 NVIDIA 研究者 2018 年的工作,并说明它被 FlashAttention 使用。

本课聚焦 max/sum 合并公式,不把所有现代 kernel 的具体访存次数都等同于这份板书。

FlashAttention 不同版本、硬件后端和编译器可能采用不同 tile、warp 归约与流水方式,但要保持精确结果,都必须等价维护相应归一化状态。

跟练与练习

原视频练习

编者练习

已扫描状态为 mold=3,old=2m_{\text{old}}=3,\ell_{\text{old}}=2,新元素为 5。更新后的 mm\ell 是什么?

查看参考答案

m=max(3,5)=5m=\max(3,5)=5。旧分母必须从基准 3 转到基准 5:
=2e35+e55=2e2+1.\ell=2e^{3-5}+e^{5-5}=2e^{-2}+1.
不能直接写成 2+12+1,否则旧项和新项不在同一指数基准。

编者练习 2

为什么在线 Softmax 的独立数组版本仍通常需要第二遍?

查看参考答案

第一次扫描结束前,最终最大值 mm 与分母 \ell 尚未确定,此前元素的最终概率无法永久写出。第二遍使用最终状态计算 exim/e^{x_i-m}/\ell。若与后续算子融合,可改为维护并重标定输出累加器,但那属于更大的融合算法。

常见误区

  • 误区:安全 Softmax 改变了概率分布。纠正:分子分母乘同一 eme^{-m},结果完全相同。
  • 误区:在线 Softmax 不再减最大值。纠正:它持续减去运行最大值,并在最大值变化时重标定旧和。
  • 误区:合并局部分母时直接做 a+b\ell_a+\ell_b。纠正:两者基准不同,必须乘相应缩放因子。
  • 误区:在线 Softmax 一遍就能写出完整概率数组。纠正:独立输出通常仍需第二遍归一化。
  • 误区:三次逻辑循环必然等于三次独立 HBM 读写。纠正:真实 kernel、cache 与编译融合会改变物理流量。
  • 误区:FlashAttention 只维护 max 与 sum。纠正:它还维护并重标定输出累加器。

本课小结

  • 朴素安全 Softmax 逻辑上依次求最大值、分母与归一化结果。
  • 在线 Softmax 用状态 (m,)(m,\ell) 在一次扫描中共同更新最大值与安全分母。
  • 最大值刷新时,旧分母必须乘 emoldmnewe^{m_{\text{old}}-m_{\text{new}}}
  • 独立 Softmax 因此从三遍降到两遍,同时保持 max-shift 数值稳定性。
  • 访存收益取决于真实融合与硬件;FlashAttention 还把在线状态扩展到输出累加。
04

主题讲解 · 03:25

FlashAttention 如何跳过因果掩码的上三角块

学习目标

  • 能从因果掩码推导上三角 Softmax 权重为零。
  • 能区分完全无效 tile、完全有效 tile 与跨越对角线的边界 tile。
  • 能解释为何跳过无效 K/V tile 同时减少计算与数据搬运。
  • 能说明局部分数 tile 为何仍需在线 Softmax 修正累加。
  • 能区分 Decoder-only 架构与单 token decode 阶段。

前置与衔接

Decoder-only Transformer 要保证位置 ii 不能看到未来位置 j>ij>i

传统写法先计算完整分数矩阵,再在上三角加 -\infty

FlashAttention 按 tile 计算,不必真的先算出那些注定被掩码为零的完整块。

核心判断是:

核心讲解

1. 因果注意力的掩码定义

忽略多头与缩放因子,原始分数为

S=QK.S=QK^\top.

因果掩码 MM 定义为

Mij={0,ji,,j>i.M_{ij}= \begin{cases} 0,&j\le i,\\ -\infty,&j>i. \end{cases}
图 1

因果注意力把位置 j>ij>i 的未来分数设为 -\infty,保证第 ii 个 query 不读取未来 token。

原视频 · 00:20 ↗

掩码后分数为

S~=S+M.\tilde S=S+M.

ii 行只有当前及历史 key 是有限值。

2. 为什么负无穷对应零贡献

逐行 Softmax 为

Pij=eS~ijkeS~ik.P_{ij} = \frac{e^{\tilde S_{ij}}} {\sum_k e^{\tilde S_{ik}}}.

对未来位置 j>ij>i

e=0,e^{-\infty}=0,

所以

Pij=0.P_{ij}=0.
图 2

逐行 Softmax 后,上三角的 -\infty 位置权重严格为 0,不会对加权和 PVPV 作贡献。

原视频 · 01:00 ↗

随后输出

Oi=jPijVjO_i=\sum_j P_{ij}V_j

中,这些项都是 0Vj0\cdot V_j,不会贡献任何值。

3. 从元素零贡献提升到 tile 零贡献

把 query 行与 key 列都分块。

若某个块中的每一对 (i,j)(i,j) 都满足 j>ij>i,该块所有掩码分数都是 -\infty

整个 Softmax 权重 tile 都为零,进而

PtileVtile=0.P_{\text{tile}}V_{\text{tile}}=0.

所以不必先计算分数再发现它为零。

运行时可以根据 tile 的位置关系预先判定并跳过。

4. 三类因果 tile

对齐分块后,可以把 tile 分为三类:

  • 对角线左下方:全部满足 jij\le i,整块有效;
  • 严格右上方:全部满足 j>ij>i,整块无效;
  • 穿过主对角线:同时包含历史与未来位置,必须在块内应用因果掩码。

只有第二类可以整块省略。

“省略上三角计算”不等于对角线边界块里完全没有 mask 判断。

5. 一个严格的区间判定

作为编者补充,设 query tile 行区间为半开区间

i[r0,r1),i\in[r_0,r_1),

key tile 列区间为

j[c0,c1).j\in[c_0,c_1).

c0r1,c_0\ge r_1,

则最小 key 下标也大于最大 query 下标,整块完全位于未来,可直接跳过。

若区间跨过对角线,就不能用整块零替代。

6. FlashAttention 的 tile 数据路径

视频以一个 query tile Q2Q_2 为例。

线程块把 Q2Q_2 与所需的 K1K_1 搬入片上 SRAM,计算局部分数

A21=Q2K1.A_{21}=Q_2K_1^\top.
图 3

FlashAttention 以 query tile 为工作单元,把需要的 Q/K 分块搬入片上 SRAM 计算局部分数。

原视频 · 01:40 ↗

这里的 tile 可以包含多个 token,尺寸由片上容量、数据类型和 kernel 配置共同决定。

7. 局部分数不写回 HBM

A21A_{21} 在片上立即参与 Softmax 状态更新与 V1V_1 加权。

图 4

局部注意力分数块 A21A_{21} 在片上参与 Softmax 与后续乘法,使用完即丢弃,不写回 HBM。

原视频 · 02:00 ↗

使用完成后即可丢弃,不必物化到 HBM。

这与“跳过上三角”是两项相关但不同的优化:

  • 不物化 A:避免完整分数矩阵的 HBM 存储与往返;
  • 跳过无效 tile:连分数乘法与对应 K/V 搬运都不做。

8. 为什么一个局部 O 不是最终输出

当前块得到的部分结果类似

O~2(1)=P21V1.\tilde O_2^{(1)}=P_{21}V_1.

Q2Q_2 还可能关注 K2,K3,K_2,K_3,\ldots 中的合法位置。

图 5

每个分数块只覆盖部分 key,所得 AijVjA_{ij}V_j 必须结合在线 Softmax 状态修正并累加到输出 tile。

原视频 · 02:40 ↗

因此它必须与后续 K/V tile 的贡献累加。

更重要的是,每个分数块只有局部最大值和分母。

新块改变运行最大值时,旧输出累加器也要按在线 Softmax 比例重标定,才能与完整一行 Softmax 等价。

9. 为什么可以不搬未来 K/V tile

考虑最早的一组 query,例如 Q1Q_1,与很靠后的 KNK_N

若二者对应的整个分数块都位于因果上三角,那么该块最终全为零。

图 6

若某个 K/V tile 完全位于当前 query tile 的因果未来区域,其分数全为 -\infty,可直接不搬运、不计算。

原视频 · 03:00 ↗

此时可以同时省去:

  • 从 HBM 加载 KNK_N
  • 从 HBM 加载对应 VNV_N
  • 计算 Q1KNQ_1K_N^\top
  • 对该零块执行 Softmax 更新;
  • 计算零权重与 VNV_N 的乘法。

收益不只是少做乘法,也包括少搬数据。

10. 跳过不会破坏 Softmax 分母

有人会担心:不计算这些项,Softmax 分母是否会改变?

不会。

被掩码项的指数本来就是

e=0.e^{-\infty}=0.

把零项从求和中删除不改变分母,也不改变最大值,因为合法行至少包含当前位置本身。

因此跳过是代数上的精确消元,不是近似稀疏化。

11. Prefill 与单 token decode 要分开

视频 ASR 把 “Decoder-only Transformer” 识别得不清楚,不能据此把课程理解成单 token decode。

在处理完整 causal 序列的训练或 Prefill 中,QQKK 都覆盖许多位置,确实存在大块上三角区域。

而标准自回归单 token decode 只计算最新 query,它可以关注所有已有 key,通常没有同样的完整上三角矩阵可跳。

因此该优化的显著场景主要是多 query 的 causal attention。

12. 负无穷在实现中不一定真的写入

数学上用 -\infty 表示 mask 最清楚。

kernel 中可能:

  • 对边界 tile 生成布尔 mask;
  • 用足够小的有限值实现;
  • 直接让无效 lane 不参与归约;
  • 对完全无效 tile 在调度层跳过。

只要最终等价于无效位置的 Softmax 权重为零,具体表示可以不同。

13. 版本边界

视频以 FlashAttention v2 讲解外 Q tile 与因果上三角跳过。

这一“完全掩码 tile 不加载、不计算”的原则并不应理解为只有 v2 才能采用。

FlashAttention 不同版本、不同后端以及 PyTorch SDPA/Triton 等实现,可能有不同分块大小、对角线处理和调度策略。

判断某版本是否实际跳过哪些 tile,应查看目标 kernel 与运行配置。

跟练与练习

原视频练习

编者练习

query tile 覆盖行 i=4,5i=4,5,key tile 覆盖列 j=6,7j=6,7。按 j>ij>i 为未来位置,这个 tile 能否整块跳过?

查看参考答案

可以。最小 key 下标 6 仍大于最大 query 下标 5,所以四个位置组合都位于因果未来区域,Softmax 权重全为零。

编者练习 2

query tile 覆盖 i=4,5i=4,5,key tile 覆盖 j=5,6j=5,6,能否整块跳过?

查看参考答案

不能。对 i=5,j=5i=5,j=5,位置合法;而 i=4,j=5i=4,j=5j=6j=6 等位置被掩码。该 tile 跨越对角线,必须计算合法部分并在块内应用 mask。

常见误区

  • 误区:先计算上三角,再把结果清零。纠正:完全掩码 tile 可在加载和乘法前就跳过。
  • 误区:所有碰到上三角的 tile 都能整块跳过。纠正:跨对角线的边界 tile 仍含合法元素。
  • 误区:不搬 K 就仍要搬对应 V。纠正:该权重块全零时,K 与 V 的整个贡献路径都可省去。
  • 误区:跳过掩码项会改变 Softmax 分母。纠正:这些项的指数本来就是零。
  • 误区:局部 AijVjA_{ij}V_j 已是最终输出。纠正:还要跨 K/V tiles 做在线修正与累加。
  • 误区:Decoder-only 就等于单 token decode。纠正:架构名称与推理阶段不同,上三角主要出现在多 query causal 计算。

本课小结

  • 因果掩码令 j>ij>i 的分数为 -\infty,Softmax 权重为零。
  • 完全位于上三角的 tile 对输出没有贡献,可以不加载 K/V、不计算分数。
  • 跨越主对角线的 tile 仍需块内 mask,不能整体删除。
  • FlashAttention 同时避免把局部分数写回 HBM,并用在线 Softmax 修正累加部分输出。
  • 该跳过是精确消元;具体 tile 调度和实现细节随版本与后端变化。
05

主题讲解 · 02:55

FlashAttention-2 为什么把 Q 放到外层循环

学习目标

  • 能画出外层 Q、内层 K/V 的 tile 循环。
  • 能说明哪些张量常驻 HBM,哪些状态只需片上暂存。
  • 能解释同一个输出 tile 为何要跨多个 K/V tile 修正累加。
  • 能比较 FlashAttention v1 与 v2 循环顺序的 Q/O 搬运差异。
  • 能区分视频聚焦的访存理由与 FlashAttention-2 的其他优化。

前置与衔接

FlashAttention 计算的数学目标仍是

O=softmax(QK)V.O=\operatorname{softmax}(QK^\top)V.

它不把完整注意力分数矩阵写回 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。

图 1

完整 Q、K、V 与最终 O 位于 HBM;计算围绕一个 query/output tile 在 HBM 与 SRAM 间组织。

原视频 · 00:20 ↗

分块计算会把当前需要的小块搬到寄存器或共享内存等片上存储。

完整分数矩阵

A=QKA=QK^\top

则不需要持久化到 HBM。

2. 分数 tile 使用后即可丢弃

对 query block QiQ_i 与 key block KjK_j,先计算

Sij=QiKj.S_{ij}=Q_iK_j^\top.

局部 SijS_{ij} 参与在线 Softmax,再与 VjV_j 相乘。

图 2

局部分数 A21,A22A_{21},A_{22} 与对应 V 块在片上使用,分数 tile 无需物化回 HBM。

原视频 · 00:40 ↗

一旦它对在线统计量和输出累加器的贡献已经吸收,原始分数 tile 就可以释放。

因此循环顺序主要影响 Q、K、V、O 及在线状态怎样搬运,而不是 A 如何写回。

3. 固定一个 query tile

视频选 Q2Q_2 作为例子。

把它搬入片上存储后,与第一个 K block 计算

S21=Q2K1.S_{21}=Q_2K_1^\top.

再读取 V1V_1,形成第一份局部输出贡献。

之后依次处理

K2,V2;K3,V3;K_2,V_2;\quad K_3,V_3;\quad\ldots
图 3

固定 Q2Q_2 后,依次与 K1,K2,K_1,K_2,\ldots 计算 A21,A22,A_{21},A_{22},\ldots,共同覆盖完整 key 轴。

原视频 · 01:40 ↗

内层循环结束时,当前 query 行已经看过全部合法 key。

4. 为什么不能简单累加局部 Softmax

设第一个分数块的局部最大值和分母为 m1,1m_1,\ell_1,第二块为 m2,2m_2,\ell_2

若分别对两个块归一化再直接相加,会让每块权重各自和为 1,结果不等于完整行 Softmax。

必须维护运行状态

m=max(m1,m2),m=\max(m_1,m_2),
=1em1m+2em2m.\ell = \ell_1e^{m_1-m} +\ell_2e^{m_2-m}.

旧输出部分和也要按相同基准重标定。

5. 每个 K/V block 都贡献到同一个 O tile

Q2Q_2,第一块产生 S21V1S_{21}V_1,第二块产生 S22V2S_{22}V_2

图 4

A21V1A_{21}V_1A22V2A_{22}V_2 等部分结果经在线 Softmax 修正后累加到同一个 O2O_2

原视频 · 02:00 ↗

更准确的在线更新形式可写为

O~new=emoldmnewO~old+emblockmnewO~block,\tilde O^{\text{new}} = e^{m_{\text{old}}-m_{\text{new}}}\tilde O^{\text{old}} +e^{m_{\text{block}}-m_{\text{new}}}\tilde O^{\text{block}},

其中 O~\tilde O 是尚未除以最终分母的加权和。

等所有 K/V blocks 处理完,再除以最终 \ell 得到 O2O_2

6. 为什么 O2 必须跨内层循环保留

O2O_2 不是一次 S21V1S_{21}V_1 就完成的。

它与在线最大值 m2m_2、分母 2\ell_2 一起构成跨 K/V blocks 的运行状态。

如果每处理一个 K/V block 都把这组状态写回 HBM,下一轮又读回来,会产生重复访存。

更好的做法是让负责 Q2Q_2 的线程块在内层循环期间保留这些状态。

7. 外 Q 内 K/V 的片上驻留

FlashAttention-2 的抽象循环可以写成:

  1. 选择一个 query tile QiQ_i
  2. QiQ_i 放入片上存储;
  3. 初始化该 tile 的 mi,i,O~im_i,\ell_i,\tilde O_i
  4. 依次流过所有合法 Kj,VjK_j,V_j
  5. 完成后一次写回最终 OiO_i
图 5

一个线程块可让 Q2Q_2 与运行中的 O2O_2 驻留 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 复用。

但对某个 Qi,OiQ_i,O_i 而言,每个外层 K/V block 到来时,都可能需要:

  • 从 HBM 重新读入 QiQ_i
  • 读回此前的 Oi,mi,iO_i,m_i,\ell_i
  • 更新后再写回。

于是 Q 与输出状态可能随 K/V 外层循环反复搬运。

9. v2 交换循环后的核心收益

交换为外层 query 后,一个线程块从头到尾负责一个 Qi/OiQ_i/O_i tile。

图 6

FlashAttention-2 采用外层遍历 query tiles、内层遍历 key/value tiles 的组织,让当前 Q/O tile 固定、K/V 块移动。

原视频 · 02:40 ↗

它只需在开始时读取 QiQ_i,在结束时写回一次最终 OiO_i

在线 Softmax 状态也可在片上存活整个内层循环。

因此减少了 Q、O 以及 m,m,\ell 的 HBM 往返。

10. 交换循环不改变数学结果

对固定 query block,完整输出是对所有 key blocks 的贡献归约。

只要在线 Softmax 的重标定正确,按 K1,K2,K_1,K_2,\ldots 的顺序逐块合并,与一次计算完整 key 轴等价。

循环重排改变数据驻留与并行任务划分,不改变最终

softmax(QiK)V.\operatorname{softmax}(Q_iK^\top)V.

有限精度下归约顺序可能带来小幅舍入差异,但不是算法近似。

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 彼此输出独立,可以由不同线程块并行处理。

每个线程块只需维护自己的 Oi,mi,iO_i,m_i,\ell_i

这为较长序列提供 query 维的并行任务。

FlashAttention-2 的论文与实现还包含减少非矩阵乘 FLOPs、改善线程块/warp 工作划分等优化。

视频本课只聚焦循环交换与 Q/O 搬运,不应把 v2 的全部收益都归因于这一点。

13. 与因果上三角跳过的结合

在 causal attention 中,固定 QiQ_i 后,内层只需遍历合法 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

为什么 S21V1+S22V2S_{21}V_1+S_{22}V_2 不能直接当作 Softmax attention 的输出?

查看参考答案

S21,S22S_{21},S_{22} 是原始分数而非完整行归一化权重。即使各块先做局部 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 细节。
06

主题讲解 · 02:32

物化注意力分数矩阵为何如此昂贵

学习目标

  • 能从 Q、K 的 shape 推导 N×NN\times N 分数矩阵。
  • 能区分平方中间存储与稠密注意力的平方计算量。
  • 能解释 FlashAttention 如何避免把完整 A 写回 HBM。
  • 能说明 O(Nd)O(Nd) 关于 NN 为线性的成立条件。
  • 能区分前向、反向与整个模型显存的不同统计口径。

前置与衔接

FlashAttention 经常被概括为把注意力显存从平方级降到线性级。

这句话容易被误解成“注意力不再计算 N2N^2 个 query-key 关系”。

视频抓住了更准确的瓶颈:

数学输出仍是精确的 dense attention,主要改变的是内存复杂度与数据搬运。

核心讲解

1. 先固定 N 与 d 的含义

对单个 attention head,设:

  • NN:序列 token 数;
  • dd:该 head 的维度。

Q,K,VRN×d.Q,K,V\in\mathbb{R}^{N\times d}.
图 1

单个 attention head 中,QQKK 的 shape 为 N×dN\times d,视频以 N104N\approx10^4d128d\approx128 说明两轴量级差异。

原视频 · 00:40 ↗

视频用长上下文 N104N\approx10^4、head dimension d128d\approx128 作量级示例。

具体模型的 dd 可以不同,但长序列时通常关注 NdN\gg d 的情形。

2. 为什么 AA 会变成 N×NN\times N

转置后

KRd×N.K^\top\in\mathbb{R}^{d\times N}.

因此

A=QKRN×N.A=QK^\top \in\mathbb{R}^{N\times N}.
图 2

N×dN\times d 的 Q 与 d×Nd\times NKK^\top 相乘,输出分数矩阵沿 query、key 两轴扩展为 N×NN\times N

原视频 · 01:00 ↗

第一根 N 轴枚举 query token,第二根 N 轴枚举 key token。

每个元素

Aij=qikjA_{ij}=q_i^\top k_j

表示一对 token 的注意力分数。

3. 平方存储从哪里来

若传统实现把 A 写回 HBM,它需要保存

N2N^2

个元素。

图 3

传统实现若把 A=QKA=QK^\topN×NN\times N 分数矩阵写回 HBM,会引入关于序列长度的 O(N2)O(N^2) 存储。

原视频 · 00:20 ↗

而单个 Q、K、V 或 O 只包含

NdNd

个元素。

两者比值为

N2Nd=Nd.\frac{N^2}{Nd}=\frac{N}{d}.

NdN\gg d 时,A 很快成为更大的中间张量。

4. 用 N=10000、d=128 感受量级

单头 Q 的元素数为

10000×128=1.28×106.10000\times128=1.28\times10^6.

A 的元素数为

100002=108.10000^2=10^8.

A 约为单个 Q 的 78 倍。

若仅为说明量级,按 FP16/BF16 每元素 2 字节,单头 A 约 200 MB。

实际总量还会乘 batch、head 数,并受 dtype、是否保存 logits/probabilities、mask 与框架实现影响。

这个数字只是 shape 估算,不是任意模型运行时的固定显存值。

5. 不好不只在“占着显存”

把 A 写回 HBM 后,后续 Softmax 和 PVPV 还要把它读回来。

因此物化 A 同时带来:

  • 大型中间张量的峰值显存;
  • 写 A 的 HBM 流量;
  • 读 A 或 P 的 HBM 流量;
  • 额外 kernel 边界与同步机会;
  • 长序列下更强的带宽压力。

许多 attention kernel 在实际硬件上受内存 I/O 限制,而不只是算术吞吐限制。

6. FlashAttention 怎样切 tile

FlashAttention 把 Q 沿 query token 行切成 tiles,把 KK^\top 沿 key token 列切成 tiles。

图 4

FlashAttention 把 Q 沿 token 行切块、把 KK^\top 沿 key 列切块,每个 tile 通常包含多个完整 token 向量。

原视频 · 01:20 ↗

在视频的简化画法里,每个 token 向量的完整 dd 维保留在 tile 中。

一个 Br×dB_r\times d 的 Q tile 与 d×Bcd\times B_cKK^\top tile 相乘,得到

StileRBr×Bc.S_{\text{tile}}\in\mathbb{R}^{B_r\times B_c}.

Br,BcB_r,B_c 根据片上容量和 kernel 设计选择。

7. 为什么不沿 d 随便切断

每个注意力分数需要完整收缩

qikj=r=1dqirkjr.q_i^\top k_j=\sum_{r=1}^{d}q_{ir}k_{jr}.

若沿 dd 切分,也不是数学上不可能,但必须再对各段部分内积做归约。

视频强调“token 完整性”,是在当前 tile 解释中把完整 head dimension 放进计算块,避免再引入 contraction 轴的跨块合并。

不能把它理解为任何 kernel 都绝不会在硬件微操作层拆分 dd

8. A tile 只在片上短暂存在

局部分数 StileS_{\text{tile}} 产生后,立即用于更新:

  • 每行运行最大值;
  • 每行 Softmax 分母;
  • 对 V 的输出加权和。
图 5

每个小分数块只在 SRAM 中参与在线 Softmax 与 VV 加权,消费后丢弃,不形成 HBM 中的完整 N×NN\times N 张量。

原视频 · 02:00 ↗

它不写回 HBM。

随后下一个 K/V tile 覆盖另一段 key 轴,在线 Softmax 公式把新旧局部状态精确合并。

所有 key blocks 处理完后,只写回最终 O tile。

9. 为什么结果仍等于完整 Softmax

分块只改变求值顺序。

对每个 query 行,FlashAttention 维护全局等价状态

mi=maxjAij,m_i=\max_j A_{ij},
i=jeAijmi,\ell_i=\sum_j e^{A_{ij}-m_i},

以及未归一化输出和

O~i=jeAijmiVj.\tilde O_i=\sum_j e^{A_{ij}-m_i}V_j.

每来一个新 tile,若最大值变化,就重标定旧 i\ell_iO~i\tilde O_i

最终

Oi=O~i/i,O_i=\tilde O_i/\ell_i,

与完整矩阵 Softmax 相同。

10. O(Nd)O(Nd) 到底指什么

不物化 A 后,主要持久 attention 张量 Q、K、V、O 都是 N×dN\times d 量级。

图 6

不物化 A 后,主要持久张量 Q、K、V、O 均为 N×dN\times d 量级;当 d 视为固定时,额外存储关于 N 为线性。

原视频 · 02:20 ↗

所以 attention 的额外持久存储可写作

O(Nd).O(Nd).

dd 对所研究的序列长度 N 视为固定模型参数时,它关于 N 是线性的。

更严谨的说法是“从 O(N2)O(N^2) 中间存储降到 O(Nd)O(Nd)”,而不是无条件把两个变量都省略成 O(N)O(N)

11. 计算复杂度仍然是平方级

Dense attention 仍需覆盖所有合法 query-key 对。

QKQK^\top 的算术量级仍是

O(N2d),O(N^2d),

PVPV 也有同阶计算。

FlashAttention 的主要贡献是 I/O-aware:减少 HBM 流量与中间存储,从而让相同 dense attention 计算在硬件上更高效。

它没有把稠密注意力变成线性 attention 算法。

12. Causal attention 的常数会变化

因果掩码只保留下三角,合法 query-key 对约为

N(N+1)2.\frac{N(N+1)}{2}.

这能把常数约减半,并允许跳过完全上三角 tiles。

但渐近计算量仍是 O(N2d)O(N^2d)

因此“跳过上三角”和“不物化完整 A”应分别理解:前者减少无效计算,后者改变中间存储与 I/O。

13. 训练反向为何也能省存储

训练需要反向传播通过 Softmax。

传统自动微分常保存大型中间激活。

FlashAttention 的反向可以利用前向保存的较小统计量和输出,重新计算局部分数与概率,而不是保存完整 A/P。

这用额外重计算换取更低内存。

具体保存哪些统计量、dropout RNG 如何恢复,会随算法版本与框架实现变化。

14. 不要把它推广成“整个模型显存线性”

模型总显存还包括:

  • 参数、梯度和优化器状态;
  • MLP 与归一化激活;
  • KV Cache;
  • logits、loss 与其他 workspace;
  • 框架内存池和碎片。

FlashAttention 消除的是 attention 中最关键的 N×NN\times N 中间物化。

它不保证整个训练或推理进程的全部显存只剩 O(Nd)O(Nd)

15. 版本边界

视频用“FlashAttention 之前的 N2N^2 注意力”概括传统 materialize-then-softmax 路径。

现代框架即使不显式标为 FlashAttention,也可能通过 memory-efficient SDPA、Triton kernel 或其他融合实现避免完整 A。

FlashAttention v1/v2/v3 与不同硬件后端的 tile、并行和反向策略也不同。

判断当前运行是否真的物化 N×NN\times N 中间量,应检查所选 backend、shape 支持条件和 profiler,而不是只看 API 名称。

跟练与练习

原视频练习

编者练习

N=8192,d=128N=8192,d=128 时,A 与单个 Q 各有多少元素?A 是 Q 的多少倍?

查看参考答案

AA81922=67,108,8648192^2=67{,}108{,}864 个元素;Q 有 8192×128=1,048,5768192\times128=1{,}048{,}576 个元素。二者比值为 N/d=64N/d=64,所以 A 是单个 Q 的 64 倍。

编者练习 2

FlashAttention 不保存完整 A,为什么计算复杂度仍不是 O(Nd)O(Nd)

查看参考答案

不保存不等于不计算。每个 query 仍需与所有合法 key 做内积,分块只让这些分数在片上短暂产生并立即消费。因此 dense attention 的 query-key 配对数仍为平方级,算术量约 O(N2d)O(N^2d)

常见误区

  • 误区:Q,KQ,K 本身就是 N×NN\times N。纠正:单头 Q,KQ,KN×dN\times d,乘积 AA 才是 N×NN\times N
  • 误区:FlashAttention 把计算复杂度降到线性。纠正:它主要把中间存储降到 O(Nd)O(Nd),dense 计算仍为 O(N2d)O(N^2d)
  • 误区:A 在 SRAM 中完整保存。纠正:只保存当前小 tile,使用后丢弃。
  • 误区:不写 A 就得近似 Softmax。纠正:在线 max、sum 与输出重标定保持精确等价。
  • 误区:O(Nd)O(Nd) 意味着 dd 也可任意增长而仍关于所有变量线性。纠正:称关于 NN 线性时把模型维度 dd 视为固定。
  • 误区:FlashAttention 让整个模型显存都变成 O(Nd)O(Nd)。纠正:它针对 attention 的平方中间量,其他状态仍存在。

本课小结

  • 单头 Q,KQ,KN×dN\times d,而分数矩阵 AAN×NN\times N
  • 传统物化 AA 会带来 O(N2)O(N^2) 存储及大规模 HBM 读写。
  • FlashAttention 分块产生 A tile,在片上完成在线 Softmax 与 V 加权后立即丢弃。
  • 主要持久 attention 张量降为 O(Nd)O(Nd),当 dd 固定时关于 NN 线性。
  • dense attention 的计算仍为 O(N2d)O(N^2d);版本与框架后端决定实际是否物化中间矩阵。
07

主题讲解 · 01:47

长序列下注意力显存差距如何被放大

学习目标

  • 能区分序列长度 N 与单头维度 d。
  • 能从 shape 推导传统 attention 的 N2N^2 中间存储。
  • 能计算 N2N^2NdNd 的元素数比值。
  • 能解释为什么短序列下差距不明显、长序列下差距剧增。
  • 能准确说明 FlashAttention 的“线性显存”指存储而非算术量。
  • 能辨别单个中间张量、attention 子层和整个模型显存三种口径。

前置与衔接

上一课已经解释了完整注意力分数矩阵为何是 N×NN\times N

本课进一步问一个更量化的问题:

答案不是固定的“慢若干倍”。

差距随 N 增大而继续放大。

要看清这个趋势,先固定单个 attention head 的记号:

Q,K,V,ORN×d.Q,K,V,O\in\mathbb{R}^{N\times d}.

其中:

  • N 是 token 数;
  • d 是该 head 的维度;
  • 比较长上下文时,常见情形是 NdN\gg d
图 1

总览板书把传统方法的 N×NN\times N 注意力中间量与 FlashAttention 的 N×dN\times d 持久张量并列,比较关于序列长度 N 的平方与线性存储。

原视频 · 00:00 ↗

核心讲解

1. 两种增长来自不同 shape

传统实现先计算

A=QK.A=QK^\top.

因为

QRN×d,KRd×N,Q\in\mathbb{R}^{N\times d}, \qquad K^\top\in\mathbb{R}^{d\times N},

所以

ARN×N.A\in\mathbb{R}^{N\times N}.

若把 A 或由它得到的 Softmax 概率写回 HBM,单头就是 N2N^2 个元素。

视频在 00:14 开始把这一路径称为平方存储。

图 2

传统实现若物化 A 或 Softmax(A),query 与 key 两根轴都随 N 增长,因此该中间量包含 N2N^2 个元素。

原视频 · 00:40 ↗

FlashAttention 仍计算精确的 dense attention。

不同点是它按块消费分数、执行在线 Softmax 并累加输出,不让完整 A 成为 HBM 中的持久中间量。

因此主要持久张量仍是 N×dN\times d 量级。

2. 为什么短序列时看不出危险

视频在 00:24 先画了 N 很小时的情况。

若 N 与 d 同量级,N2N^2NdNd 的数值不会相差太远。

例如 N=10、d=8:

N2=100,Nd=80.N^2=100, \qquad Nd=80.

只看这个规模,完整分数矩阵似乎“不算大”。

图 3

当 N 较小时,QKQK^\top 产生的 N×NN\times N 矩阵在图形上并不比 N×dN\times d 张量大很多,平方增长的危险尚不直观。

原视频 · 00:20 ↗

这正是容易形成误判的地方:

  • d 通常由模型架构固定;
  • N 会随着上下文长度扩展;
  • 平方项同时在 query 与 key 两根轴上增长。

3. 比值不是常数,而是 N/d

把 A 与一个 N×dN\times d 张量比较:

N2Nd=Nd.\frac{N^2}{Nd}=\frac{N}{d}.

因此,只要 d 固定,N 每扩大一倍,这个相对倍数也扩大一倍。

若把 Q、K、V、O 四个 N×dN\times d 张量合计作为参照,则比值为

N24Nd=N4d.\frac{N^2}{4Nd}=\frac{N}{4d}.

这两个比值口径不同:

  • N/dN/d:A 相对单个 N×dN\times d 张量;
  • N/(4d)N/(4d):A 相对 Q、K、V、O 元素数总和。

不能把二者混用。

4. 用视频的长序列示例计算

视频在 00:59 取 N=100000,且在 01:07 说明 d 可取 128 作示意。

图 4

长序列示例中 NdN\gg d;板书以 N105N\approx10^5d128d\approx128 强调 N2N^2NdNd 的量级差会迅速拉大。

原视频 · 01:00 ↗

一个 N×dN\times d 张量包含

100000×128=12,800,000100000\times128 =12{,}800{,}000

个元素。

完整 A 则包含

1000002=10,000,000,000100000^2 =10{,}000{,}000{,}000

个元素。

A 相对单个 N×dN\times d 张量大

100000128=781.25\frac{100000}{128}=781.25

倍。

相对四个 Q、K、V、O 的元素数总和,仍约大

1000004×128195.31\frac{100000}{4\times128}\approx195.31

倍。

若只作元素字节的理想化估算,FP16 的 A 单头就需要约

1010×2 bytes=20 GB.10^{10}\times2\text{ bytes}=20\text{ GB}.

这个数字只用于说明平方项规模;真实实现还受 batch、head 数、mask、重计算、融合策略和数据类型影响。

5. FlashAttention 省掉的到底是什么

视频在 00:50 明确指出 FlashAttention 不写回 A。

01:27 又用长序列图重申这一点。

图 5

FlashAttention 不把完整 N×NN\times N 注意力矩阵写回 HBM,主要持久张量保持在 N×dN\times d 量级;这里说的是存储而非稠密 attention 的算术量。

原视频 · 01:20 ↗

所以“线性”更准确地说是:

它不意味着:

  • 稠密 attention 的 query-key 配对数变成 O(N)O(N)
  • 整个模型训练显存只剩 O(Nd)O(Nd)
  • 所有实现的常数项都相同;
  • 长序列没有计算代价。

标准 dense attention 的核心算术量仍约为 O(N2d)O(N^2d)

跟练与练习

编者练习

把平方项换算成倍数 单个 head 取 N=32768、d=128。

  1. A 有多少个元素?
  2. 一个 N×dN\times d 张量有多少个元素?
  3. A 是单个 N×dN\times d 张量的多少倍?
  4. 若两者均为 FP16,A 理想化占多少 GiB?
查看参考答案

N2=327682=1,073,741,824.N^2=32768^2=1{,}073{,}741{,}824.
Nd=32768×128=4,194,304.Nd=32768\times128=4{,}194{,}304.
比值为
Nd=32768128=256.\frac{N}{d}=\frac{32768}{128}=256.
FP16 按每元素 2 bytes 估算:
1,073,741,824×2=2,147,483,648 bytes=2 GiB.1{,}073{,}741{,}824\times2=2{,}147{,}483{,}648\text{ bytes}=2\text{ GiB}.
这里只计算一个单头 A,不代表完整训练显存。

快速判断

判断下列说法:

  • “N 固定时,增大 d 会让 A 本身变大。”
  • “d 固定时,N 翻倍会让 A 元素数变为四倍。”
  • “FlashAttention 不存完整 A,所以不再计算全部 dense attention 关系。”

答案依次是:错、对、错。

常见误区

误区 1:线性比平方永远固定快某个倍数

错。

相对倍数含有 N/dN/d,会随 N 变化。

误区 2:只写 O(N)O(N),忽略 d

更严格的 shape 口径是 O(Nd)O(Nd)

只有把 d 视为固定模型常数时,才简称“关于 N 线性”。

误区 3:把内存复杂度当成计算复杂度

FlashAttention 的主要贡献是 IO-aware 分块与不物化平方中间量。

标准稠密 attention 的算术关系仍是平方级。

误区 4:把 A 的大小当成全部显存

完整训练还包括参数、优化器状态、其他层激活、梯度、临时 workspace 等。

本课只比较 attention 路径中的关键中间量。

误区 5:看到 FP16 估算就认为实现一定分配同样字节

真实 kernel 可能重计算、融合、使用不同累加精度,也可能受布局与对齐影响。

元素数分析用于判断渐近主项,不替代 profiler。

本课小结

  • 传统物化注意力矩阵需要 N2N^2 个元素。
  • FlashAttention 主要持久张量处于 NdNd 量级。
  • A 相对单个 N×dN\times d 张量的元素数比值是 N/dN/d
  • N 很小时差距不明显,N 远大于 d 时差距快速放大。
  • N=100000、d=128 时,A 相对单个 N×dN\times d 张量约大 781 倍。
  • “线性显存”是关于 N 的存储结论,不是 dense attention 的线性计算结论。
08

主题讲解 · 03:07

在线 Softmax 与 FlashAttention 的状态量

学习目标

  • 能说明在线 Softmax 每行维护的两个核心状态。
  • 能说明 FlashAttention 前向为何还需要输出累加状态。
  • 能推导跨块最大值、分母和输出分子的重标定公式。
  • 能区分未归一化输出累加器与已归一化输出的两种记法。
  • 能解释“动态规划量”是教学类比,而非传统 DP 表格。
  • 能说明不同 FlashAttention 版本与 kernel 实现的存储细节边界。

前置与衔接

普通安全 Softmax 对一行 logits x1,,xnx_1,\ldots,x_n 先取

m=maxixi,m=\max_i x_i,

再计算

=iexim.\ell=\sum_i e^{x_i-m}.

概率为

pi=exim.p_i=\frac{e^{x_i-m}}{\ell}.

在线 Softmax 的关键是:不必一次看到整行,仍能维护出同一个 m 与 \ell

视频在 00:06 特别说明,把这看成“动态规划”主要是帮助理解局部状态如何维护全局结果。

它并不是经典的二维 DP 表格。

核心讲解

1. 在线 Softmax 的两个状态

视频在 00:12 给出数量结论:

  • 在线 Softmax:两个核心状态;
  • FlashAttention 前向:三个概念状态。
图 1

在线 Softmax 每行维护运行最大值 m 与安全分母 \ell;FlashAttention 前向还要维护输出累加状态,概念上共有三类状态量。

原视频 · 00:20 ↗

对每一行而言,在线 Softmax 的状态是:

  1. 运行最大值 m;
  2. 以 m 为指数基准的安全分母 \ell

这里的 \ell 是标量。

m 也是标量。

它们都是“每行一个”,不是整个 attention 矩阵只保存一个。

2. 两个块如何合并

设旧状态覆盖集合 A:

mA=maxiAxi,m_A=\max_{i\in A}x_i,
A=iAeximA.\ell_A=\sum_{i\in A}e^{x_i-m_A}.

新块 B 的局部状态为

mB=maxiBxi,m_B=\max_{i\in B}x_i,
B=iBeximB.\ell_B=\sum_{i\in B}e^{x_i-m_B}.

合并后的最大值是

m=max(mA,mB).m=\max(m_A,m_B).

A\ell_AB\ell_B 使用了不同指数基准,不能直接相加。

先统一到新的 m:

=emAmA+emBmB.\ell =e^{m_A-m}\ell_A +e^{m_B-m}\ell_B.
图 2

分块安全 Softmax 先以各块局部最大值为基准计算指数和;合并时必须把旧分母重标定到新的全局最大值。

原视频 · 00:40 ↗

视频从 00:45 开始解释这个分母合并。

每个缩放因子都不大于 1,因此保留了安全 Softmax 的数值稳定性。

3. 为什么局部概率不能直接拼接

视频用 [1,2][1,2][3,5][3,5] 两块演示。

第一块局部最大值是 2,第二块局部最大值是 5。

于是局部指数分别是:

[e1,e0],[e^{-1},e^0],

[e2,e0].[e^{-2},e^0].
图 3

两块 logits 各自减去局部最大值后得到局部指数权重;这些权重不能直接拼接,必须先统一指数基准。

原视频 · 01:20 ↗

第二块的 e0e^0 与第一块的 e0e^0 并不代表相同的全局权重。

因为它们分别以 5 和 2 为参考。

把第一块改写到全局最大值 5 的基准,需要整体乘

e25=e3.e^{2-5}=e^{-3}.

这就是“局部信息维护全局信息”的核心。

4. FlashAttention 增加的第三类状态

attention 输出一行是

O=ipiVi.O=\sum_i p_iV_i.

若 V 也分块流入 SRAM,kernel 不能只维护 m 与 \ell

它还要维护与 V 维度相同的输出累加状态。

视频在 01:08 转入矩阵分块,在 01:47 说明 SRAM 只能容纳部分 V。

图 4

片上 SRAM 容量有限,V1,V2V_1,V_2 分块流入并被替换;因此输出 O 也要随每个块在线更新,而不能等完整概率矩阵出现。

原视频 · 02:00 ↗

因此,“三个量”应理解为三类每行状态:

  • 标量最大值 m;
  • 标量安全分母 \ell
  • 向量输出累加器。

第三项不是一个标量。

5. 用未归一化输出分子推导

最清楚的记法是定义块 B 的未归一化输出分子

O~B=iBeximBVi.\widetilde O_B =\sum_{i\in B}e^{x_i-m_B}V_i.

旧累加器为

O~A=iAeximAVi.\widetilde O_A =\sum_{i\in A}e^{x_i-m_A}V_i.

统一到新的全局最大值 m 后:

O~=emAmO~A+emBmO~B.\widetilde O =e^{m_A-m}\widetilde O_A +e^{m_B-m}\widetilde O_B.

最终归一化输出是

O=O~.O=\frac{\widetilde O}{\ell}.
图 5

各块未归一化输出需乘 embme^{m_b-m} 重标定后求和,最后再除以全局分母 \ell;这才与整行 Softmax 加权结果等价。

原视频 · 02:40 ↗

视频从 02:28 开始强调局部输出还需要按最大值差缩放,并在 02:39 补上最终除以全局分母。

6. 若实现维护的是已归一化 O

有些讲解或实现把第三个状态直接记成已归一化输出 O。

OA=O~AA,OB=O~BB.O_A=\frac{\widetilde O_A}{\ell_A}, \qquad O_B=\frac{\widetilde O_B}{\ell_B}.

则更新式要写成

O=emAmAOA+emBmBOB.O =\frac{ e^{m_A-m}\ell_AO_A +e^{m_B-m}\ell_BO_B }{\ell}.

这与未归一化记法完全等价。

但不能把已归一化 O 直接乘指数因子后相加,再忘掉 A\ell_AB\ell_B

7. “三个状态”不是所有实现字段的总数

三类状态是理解前向数学合并所需的最小概念集合。

具体 kernel 还可能保存:

  • log-sum-exp;
  • dropout 随机数状态;
  • mask 或序列边界信息;
  • tile 索引与流水线元数据;
  • 反向阶段所需的辅助量。

FlashAttention-1、FlashAttention-2、FlashAttention-3 及不同框架封装的具体寄存器、SRAM 与 HBM 布局并不相同。

所以本课的“2 个/3 个”是算法状态分类,不是某个版本源码中变量名的机械计数。

跟练与练习

编者练习

合并两块 Softmax 状态 给定: (mA,A,O~A)=(2,1+e1,e1V1+V2),(m_A,\ell_A,\widetilde O_A)=(2,1+e^{-1},\,e^{-1}V_1+V_2), (mB,B,O~B)=(5,1+e2,e2V3+V4).(m_B,\ell_B,\widetilde O_B)=(5,1+e^{-2},\,e^{-2}V_3+V_4). 写出合并后的 m、\ellO~\widetilde O 与 O。

查看参考答案

全局最大值为
m=5.m=5.
第一块需乘 e25=e3e^{2-5}=e^{-3},第二块基准不变:
=e3(1+e1)+(1+e2).\ell =e^{-3}(1+e^{-1})+(1+e^{-2}).
未归一化输出分子为
O~=e3(e1V1+V2)+(e2V3+V4).\widetilde O =e^{-3}(e^{-1}V_1+V_2) +(e^{-2}V_3+V_4).
最后
O=O~.O=\frac{\widetilde O}{\ell}.
展开可验证它等于以全局最大值 5 计算的整行 attention 输出。

自检:状态的 shape

若一个 query block 有 BrB_r 行,head dimension 为 d,则:

  • m 的 shape 是 BrB_r
  • \ell 的 shape 是 BrB_r
  • O~\widetilde O 的 shape 是 Br×dB_r\times d

这比“总共三个标量”更准确。

常见误区

误区 1:局部 Softmax 概率可以直接拼接

错。

每块的局部最大值与局部分母不同,必须统一基准并重新归一化。

误区 2:在线 Softmax 只维护最大值

错。

最大值保证指数稳定,但还需要分母才能恢复正确概率。

误区 3:FlashAttention 的第三个状态是最终 O

不一定。

数学上既可维护未归一化 O~\widetilde O,也可维护已归一化 O;两种更新式不同。

误区 4:“动态规划量”就是 DP 数组

本课使用的是类比。

更准确地说,这些是可合并的充分状态或归约状态。

误区 5:不同版本一定保存同样的中间量

错。

版本、精度、mask、dropout 与硬件都会改变实现细节,但最大值重标定的数学原则不变。

本课小结

  • 在线 Softmax 每行维护运行最大值 m 与安全分母 \ell
  • 跨块合并时,局部分母必须乘 embme^{m_b-m} 后求和。
  • FlashAttention 前向还要维护输出累加状态,因此概念上有三类状态。
  • 未归一化输出分子与分母使用完全相同的最大值重标定因子。
  • 最终输出由 O=O~/O=\widetilde O/\ell 得到。
  • “2 个/3 个”是算法状态分类,不是某个 kernel 的全部实现变量。
09

主题讲解 · 03:23

用矩阵链式法则推导注意力反向梯度

学习目标

  • 能从标量乘积链式法则过渡到矩阵乘积反向传播。
  • 能使用 Frobenius 内积检查矩阵梯度的方向。
  • 能推导 S=QKS=QK^\top 对 Q、K 的梯度。
  • 能推导 O=PVO=PV 对 P、V 的梯度。
  • 能写出逐行 Softmax 的 Jacobian-vector product。
  • 能说明缩放、mask、dropout 等实际 attention 细节应插入何处。

前置与衔接

视频试图用一句口诀快速记忆矩阵乘法反向:

这个口诀有用,但必须补上两个条件:

  1. 矩阵乘法不可交换,因子次序不能随意改;
  2. 最终公式要通过 shape 或微分内积检查。

视频在 00:14 从一个二级结论切入,在 00:27 先回顾标量乘积。

图 1

从标量乘积 y=abcdy=abcd 出发,对 b 求导时把上游梯度放到 b 的位置,并乘上其余因子,作为矩阵推广的直觉起点。

原视频 · 00:20 ↗

核心讲解

1. 标量乘积为何可以“去掉被求导因子”

y=abcd,y=abcd,

损失 L 通过 y 依赖 b。

Lb=Lyyb=Lyacd.\frac{\partial L}{\partial b} =\frac{\partial L}{\partial y} \frac{\partial y}{\partial b} =\frac{\partial L}{\partial y}acd.

标量乘法可交换,因此看起来像“把 b 替换成上游梯度,剩下的照乘”。

矩阵乘法没有交换律,所以推广时不能只照搬“剩下的照乘”。

2. 矩阵乘积的可靠推导工具

Y=ABCD,Y=ABCD,

上游梯度记为

G=LY.G=\frac{\partial L}{\partial Y}.

只让 B 变化:

dY=A(dB)CD.dY=A\,(dB)\,CD.

用 Frobenius 内积定义梯度:

dL=G,dYF=tr(GdY).dL=\langle G,dY\rangle_F =\operatorname{tr}(G^\top dY).

代入并循环移动 trace 中的因子:

dL=tr ⁣(GA(dB)CD)=tr ⁣((AGDC)dB).dL =\operatorname{tr}\!\left(G^\top A(dB)CD\right) =\operatorname{tr}\!\left((A^\top G D^\top C^\top)^\top dB\right).

因此

LB=AGDC.\frac{\partial L}{\partial B} =A^\top G D^\top C^\top.
图 2

矩阵乘积的反向传播不能任意交换因子;对中间矩阵求梯度时要保持乘法次序,并用转置使 shape 对齐。

原视频 · 01:20 ↗

视频在 01:08 开始把标量结论推广到矩阵。

板书的紧凑写法 AG(CD)A^\top G(CD)^\top 与上式相同,因为

(CD)=DC.(CD)^\top=D^\top C^\top.

3. shape 检查比口诀更可靠

假设

ARa×b,BRb×c,CRc×d,DRd×e.A\in\mathbb{R}^{a\times b}, \quad B\in\mathbb{R}^{b\times c}, \quad C\in\mathbb{R}^{c\times d}, \quad D\in\mathbb{R}^{d\times e}.

Y,GRa×e.Y,G\in\mathbb{R}^{a\times e}.

候选梯度

AGDCA^\top G D^\top C^\top

的 shape 是

(b×a)(a×e)(e×d)(d×c)=b×c,(b\times a)(a\times e)(e\times d)(d\times c) =b\times c,

正好等于 B 的 shape。

若 shape 不对,转置或次序一定有错。

4. 注意力前向的最简计算图

本课采用视频中的简化记号:

S=QK,S=QK^\top,
P=softmax(S),P=\operatorname{softmax}(S),
O=PV.O=PV.

反向从上游 dO 开始,依次得到:

dP,dV,dS,dQ,dK.dP,dV,dS,dQ,dK.

这就是板书左侧所列的五条主要梯度支路。

5. 由 O=PV 推出 dP 与 dV

微分为

dO=(dP)V+P(dV).dO=(dP)V+P(dV).

分别收集两支:

dP=dOV,dP=dO\,V^\top,
dV=PdO.dV=P^\top dO.
图 3

O=PVO=PV,反向传播得到 dP=dOVdP=dO V^\topdV=PdOdV=P^\top dO,shape 检查可快速发现转置错误。

原视频 · 02:20 ↗

视频在 02:14 开始把 P、V 两支与 Q、K 两支类比。

shape 检查:若

PRNq×Nk,VRNk×dv,ORNq×dv,P\in\mathbb{R}^{N_q\times N_k}, \quad V\in\mathbb{R}^{N_k\times d_v}, \quad O\in\mathbb{R}^{N_q\times d_v},

dOV:(Nq×dv)(dv×Nk)=Nq×Nk,dO\,V^\top: (N_q\times d_v)(d_v\times N_k) =N_q\times N_k,

与 P 相同。

PdO:(Nk×Nq)(Nq×dv)=Nk×dv,P^\top dO: (N_k\times N_q)(N_q\times d_v) =N_k\times d_v,

与 V 相同。

6. 由 S=QK^T 推出 dQ 与 dK

微分为

dS=(dQ)K+Q(dK).dS=(dQ)K^\top+Q(dK)^\top.

对应梯度:

dQ=dSK,dQ=dS\,K,
dK=dSQ.dK=dS^\top Q.
图 4

S=QKS=QK^\top,上游梯度为 dS 时有 dQ=dSKdQ=dS KdK=dSQdK=dS^\top Q;K 原本带转置,因此两支公式不完全对称。

原视频 · 02:00 ↗

视频在 01:36 进入 dQ,在 02:27 进入较容易出错的 dK。

dK 之所以看起来多一步,是因为前向中出现的是 KK^\top

可先求

d(K)=QdS,d(K^\top)=Q^\top dS,

再转置:

dK=(QdS)=dSQ.dK=(Q^\top dS)^\top=dS^\top Q.

7. Softmax 反向不能套矩阵乘法口诀

视频在 03:04 转到 dS,并在 03:11 用一个抽象的 Softmax 反向算子带过细节。

图 5

P=softmax(S)P=\operatorname{softmax}(S) 的反向不是普通矩阵乘法;dS 需由 Softmax 的 Jacobian-vector product 从 dP 与 P 计算。

原视频 · 03:00 ↗

对单行

p=softmax(s),p=\operatorname{softmax}(s),

其 Jacobian 为

J=diag(p)pp.J=\operatorname{diag}(p)-pp^\top.

给定上游梯度 g=dP,该行的 dS 为

dS=Jg=p(gg,p1).dS =J^\top g =p\odot\left(g-\langle g,p\rangle\mathbf 1\right).

矩阵逐行写成

dS=P(dProwsum(dPP)),dS =P\odot \left( dP-\operatorname{rowsum}(dP\odot P) \right),

其中 rowsum 的结果在该行广播。

8. 实际 attention 中还缺哪些因子

标准 scaled dot-product attention 常写为

S=QKdk+M.S=\frac{QK^\top}{\sqrt{d_k}}+M.

若缩放因子

c=1/dk,c=1/\sqrt{d_k},

dQ=cdSK,dK=cdSQ.dQ=c\,dS K, \qquad dK=c\,dS^\top Q.

视频板书为突出矩阵链式法则,省略了这个缩放。

mask M 若为常量,不需要对 M 求梯度,但被屏蔽位置的概率与梯度处理必须符合具体实现。

训练中的 dropout、causal mask、GQA/MQA、变长序列也会改变 kernel 路径,但不改变上述局部矩阵微分原则。

9. 这套推导与 FlashAttention 的关系

这些是数学上的反向公式。

FlashAttention 的工程难点不是改变梯度定义,而是:

  • 按块重计算前向需要的 S、P;
  • 避免完整 N×NN\times N 中间量写回 HBM;
  • 在片上内存有限条件下组织 dQ、dK、dV 的归约;
  • 控制数值误差与数据搬运。

因此,“能写出五条梯度公式”不等于“已经实现 FlashAttention backward kernel”。

跟练与练习

编者练习

只用 shape 推出两条梯度 给定 S=QK,S=QK^\top, 其中 QRNq×d,KRNk×d,dSRNq×Nk.Q\in\mathbb{R}^{N_q\times d}, \quad K\in\mathbb{R}^{N_k\times d}, \quad dS\in\mathbb{R}^{N_q\times N_k}. 写出 dQ、dK,并验证 shape。

查看参考答案

dQ=dSK.dQ=dS K.
shape:
(Nq×Nk)(Nk×d)=Nq×d.(N_q\times N_k)(N_k\times d) =N_q\times d.
dK=dSQ.dK=dS^\top Q.
shape:
(Nk×Nq)(Nq×d)=Nk×d.(N_k\times N_q)(N_q\times d) =N_k\times d.
两者分别与 Q、K 的 shape 相同。

进阶自检

若前向为

S=cQKS=cQK^\top

且 c 是常量,dQ 与 dK 都应额外乘 c。

若遗漏 c,数值梯度检查会失败。

常见误区

误区 1:矩阵乘法像标量一样可以交换

错。

转置会反转乘积次序,shape 是最直接的检查手段。

误区 2:“对谁求导就替换谁”足以覆盖所有算子

它只适合作为矩阵乘法的记忆辅助。

Softmax、mask、dropout 等需要各自的局部反向规则。

误区 3:dK=QdSdK=Q^\top dS

QdSQ^\top dS 是对 KK^\top 的梯度 shape。

对 K 的梯度还需转置,得到 dSQdS^\top Q

误区 4:板书省略缩放就代表实际 attention 没有缩放

本课使用的是简化计算图。

实现时要把 1/dk1/\sqrt{d_k} 放回 dQ、dK。

误区 5:推导数学梯度等于解决内存问题

数学公式定义结果。

FlashAttention backward 还要决定哪些量重计算、如何分块以及如何归约。

本课小结

  • 标量乘积的直觉可以推广,但矩阵因子不能交换。
  • Frobenius 内积与 trace 是可靠的矩阵梯度推导工具。
  • O=PVO=PV,有 dP=dOVdP=dOV^\topdV=PdOdV=P^\top dO
  • S=QKS=QK^\top,有 dQ=dSKdQ=dSKdK=dSQdK=dS^\top Q
  • Softmax 反向是逐行 Jacobian-vector product,不是普通矩阵乘法。
  • 实际 scaled attention 还要把缩放、mask、dropout 与版本细节纳入计算图。
10

主题讲解 · 03:20

FlashAttention 反向传播为何仍是线性额外存储

学习目标

  • 能区分前向注意力矩阵与反向注意力梯度矩阵。
  • 能解释传统训练为何可能持久化平方级中间量。
  • 能说明 FlashAttention backward 如何用分块重计算换存储。
  • 能从 dQ、dK、dV 的 shape 判断主要持久梯度规模。
  • 能准确限定“关于 N 线性”的假设与统计范围。
  • 能说明不同 FlashAttention 版本和框架实现不必采用完全相同的缓存策略。

前置与衔接

前向阶段的主要问题是完整注意力分数或概率矩阵:

S,PRN×N.S,P\in\mathbb{R}^{N\times N}.

反向阶段又会出现同 shape 的上游中间梯度:

dS,dPRN×N.dS,dP\in\mathbb{R}^{N\times N}.

因此,只省掉前向 P 还不够。

如果 backward 又把 dS 或 dP 物化回 HBM,平方存储会回来。

视频在 00:30 先给结论:FlashAttention 的前向与反向都能避免平方级持久中间量。

图 1

板书先给结论:前向不把完整注意力矩阵写回 HBM,反向也不把其 N×NN\times N 梯度矩阵写回 HBM。

原视频 · 00:20 ↗

核心讲解

1. “线性”首先是关于 N 的说法

视频在 00:05 明确把线性限定为序列长度 N。

单个 head 中:

Q,K,V,ORN×d.Q,K,V,O\in\mathbb{R}^{N\times d}.

相应梯度:

dQ,dK,dV,dORN×d.dQ,dK,dV,dO\in\mathbb{R}^{N\times d}.

当 d 视为固定模型维度时,这些张量关于 N 线性。

S,P,dS,dPRN×NS,P,dS,dP\in\mathbb{R}^{N\times N}

关于 N 平方。

“线性显存”更严格地应写为 O(Nd)O(Nd),不是无条件的 O(N)O(N)

2. 传统训练路径为什么有平方中间量

简化前向计算图为

S=QK,S=QK^\top,
P=softmax(S),P=\operatorname{softmax}(S),
O=PV.O=PV.

朴素自动微分若为 backward 保存前向激活,可能把 S 或 P 写回 HBM。

反向又计算

dP=dOV,dP=dOV^\top,

dS=softmaxBackward(dP,P).dS=\operatorname{softmaxBackward}(dP,P).

dP、dS 也是 N×NN\times N

图 2

传统训练路径会持久化 S/P 等 N×NN\times N 中间量及相关梯度,因而平方项可能主导 attention 子层的激活存储。

原视频 · 02:40 ↗

视频在 02:25 区分 S 与 P,并在 02:39 开始比较传统训练的写回。

3. backward 并不需要把 dP、dS 全部存下来

最终需要交给上一层的是:

dQ,dK,dV.dQ, \quad dK, \quad dV.

它们都是 N×dN\times d

dP 与 dS 只是计算这些最终梯度的中间量。

FlashAttention backward 可以对 Q、K、V 分块:

  1. 重新读取某个 Q、K tile;
  2. 在片上重算局部分数 S tile;
  3. 利用前向保存的每行归一化信息重构局部 P tile;
  4. 计算局部 dP、dS;
  5. 立即把它们用于累加 dQ、dK、dV;
  6. 用完后丢弃局部平方 tile。

因此它计算过平方数量的局部元素,却不让完整 N×NN\times N 梯度成为持久 HBM 张量。

图 3

FlashAttention 的关键不只是省掉前向 A,还通过重计算与分块反向避免持久化 L/A\partial L/\partial A

原视频 · 00:40 ↗

视频在 00:38 说明前向不写回注意力矩阵,在 00:44 说明反向不写回其梯度矩阵。

4. 用梯度公式看最终持久张量

S=QKS=QK^\top

得到

dQ=dSK,dQ=dS K,
dK=dSQ.dK=dS^\top Q.

O=PVO=PV

得到

dV=PdO.dV=P^\top dO.
图 4

S=QKS=QK^\topdQ=dSKdQ=dS K,由 O=PVO=PVdV=PdOdV=P^\top dO;两式展示反向矩阵乘法的 shape 对齐。

原视频 · 02:00 ↗

视频在 01:34 从 Q 支路说明 dQ,在 02:02 转到 dV。

这些公式看起来引用完整 dS 或 P。

但矩阵乘法可以按 tile 做,并把局部贡献规约到最终 dQ、dK、dV。

数学公式定义结果,分块顺序决定内存行为。

5. 反向线性存储的代价是重计算

若前向不保存 P,backward 仍需要 P 才能求 Softmax 梯度。

常见策略是保存每行 log-sum-exp 或等价的归一化统计量,然后重算局部

Sij=qikjS_{ij}=q_i^\top k_j

Pij=eSijLSEi.P_{ij}=e^{S_{ij}-\operatorname{LSE}_i}.

这体现典型的时间—空间权衡:

  • 少存平方激活;
  • backward 多做一部分局部前向重计算;
  • 换来显著减少 HBM 读写与容量压力。

“重计算增加 FLOPs”不必然意味着墙钟时间更差,因为 FlashAttention 的目标正是减少高成本 IO。

实际收益依赖硬件、shape、精度与 kernel 版本。

6. 为什么主要持久状态仍是 Nd

概念上,前向主要持久输入/输出为:

Q,K,V,O=O(Nd).Q,K,V,O=O(Nd).

反向主要持久梯度为:

dQ,dK,dV,dO=O(Nd).dQ,dK,dV,dO=O(Nd).

此外还可能保存每行统计量:

O(N).O(N).

局部 S、P、dS、dP tile 的大小由块尺寸控制,不随完整 N2N^2 物化。

图 5

FlashAttention 反向以分块重计算换取不存完整 S、P、dS、dP;固定 head dimension d 时,主要持久激活与梯度关于 N 为 O(Nd)O(Nd)

原视频 · 03:00 ↗

视频在 03:00 画出 FlashAttention 只保留 Q/K/V/O 与相应梯度的对比,在 03:12 收束到 O(Nd)O(Nd)

7. 矩阵链式法则只是数学入口

视频在 00:56 回顾矩阵乘法链式法则。

图 6

矩阵乘积反向需把上游梯度接入对应位置,并按 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/\ell 还是其他辅助量;
  • backward 的并行划分与归约方式;
  • 是否启用 dropout、causal、window attention;
  • GQA/MQA 对 K/V head 的共享方式;
  • workspace、对齐与临时缓冲大小;
  • 自动微分框架外围是否保存额外张量。

“线性额外存储”也不等于“整个模型显存线性”。

参数、优化器状态、MLP 激活、残差、通信缓冲等仍在统计范围之外。

跟练与练习

编者练习

判断哪些张量必须完整持久化 单头 self-attention 中,给出 Q,K,V,dORN×d.Q,K,V,dO\in\mathbb{R}^{N\times d}. 若 backward 对局部 S、P、dP、dS 即算即用,列出:

  1. 最终必须输出的梯度;
  2. 可只以 tile 临时存在的平方中间量;
  3. 关于 N 的主要持久梯度复杂度。
查看参考答案

最终必须输出:
dQ,dK,dVRN×d.dQ,dK,dV\in\mathbb{R}^{N\times d}.
可按 tile 临时存在:
S,P,dP,dS.S,P,dP,dS.
它们在数学上覆盖 N×NN\times N 关系,但无需完整写回 HBM。
固定 d 时,最终梯度与每行辅助量的主要持久规模为
O(Nd)+O(N)=O(Nd).O(Nd)+O(N)=O(Nd).
这不包含模型其他层、参数、优化器与框架 workspace。

思考题

为什么 backward 重算 P 不会改变数学梯度?

因为在相同 Q、K、mask、缩放与数值规则下,重算得到的是同一前向中间量;它改变保存策略,不改变计算图定义。

若 dropout 存在,则必须能重现相同 dropout mask,不能随意重新采样。

常见误区

误区 1:前向不存 P,反向就无法求梯度

可以利用保存的归一化统计量按块重算 P。

误区 2:计算过 N2N^2 个梯度元素就必须存 N2N^2

错。

局部 tile 可即算即用并立刻规约。

误区 3:反向线性显存意味着反向线性计算

错。

稠密 attention 的 backward 算术量仍含平方级 token 配对。

误区 4:所有 FlashAttention 版本峰值显存完全相同

错。

渐近主项一致不代表常数、workspace 与保存字段一致。

误区 5:O(Nd)O(Nd) 就是整卡训练显存

本课只讨论 attention 算子的主要额外激活/梯度存储。

整模型还包含大量其他状态。

本课小结

  • 反向阶段的 dS、dP 也具有 N×NN\times N shape,是潜在平方存储源。
  • FlashAttention backward 按块重算 S、P、dP、dS,并立即累加到 dQ、dK、dV。
  • 完整平方中间梯度无需持久写回 HBM。
  • 最终 dQ、dK、dV 均为 N×dN\times d,固定 d 时关于 N 线性。
  • 线性额外存储通过重计算和 IO 优化换取,不代表线性算术复杂度。
  • 具体保存量与峰值显存仍受 FlashAttention 版本、框架、mask、dropout 与硬件影响。
11

主题讲解 · 02:48

在线 Softmax 的循环与规约

学习目标

  • 能准确区分局部最大值、全局最大值与运行最大值。
  • 能定义与最大值绑定的安全 Softmax 分母。
  • 能推导两个局部状态的分母重标定公式。
  • 能证明顺序循环与分块规约得到同一数学结果。
  • 能写出在线 Softmax 状态合并的单位元与结合性直觉。
  • 能区分数学等价与有限精度浮点下的逐位一致。

前置与衔接

本地 Whisper 把“局部、最大值、线程、遍历”等识别成近音字,结尾还有超出有效内容的重复短句。 本课术语与公式均按板书及上下文共同校正,不采用字幕尾部重复文本。

视频的核心不是背一个实现细节,而是理解一个可合并状态:

(m,).(m,\ell).

其中:

  • m 是当前已见 logits 的最大值;
  • \ell 是以 m 为指数基准的安全 Softmax 分母。

视频在 00:05 用“局部信息维护全局信息”概括这件事。

核心讲解

1. 循环与规约处理的是同一种状态

板书在 00:20 区分两种执行方式:

  • 单个线程内逐元素循环;
  • 线程或块之间做规约。
图 1

板书把在线 Softmax 解释为:单个线程内用循环逐步更新状态,线程或块之间再用同一合并规则做规约。

原视频 · 00:20 ↗

二者并不是两套不同的 Softmax 公式。

区别只在于数据如何分组:

  • 循环:旧状态与一个新元素合并;
  • 规约:两个已经聚合好的局部状态合并。

只要合并算子正确,两条路径就恢复同一个全局最大值和全局分母。

2. 每个局部块保存什么

对一个索引集合 A,定义

mA=maxiAxi,m_A=\max_{i\in A}x_i,

以及

A=iAeximA.\ell_A=\sum_{i\in A}e^{x_i-m_A}.

A\ell_A 不是原始指数和

iAexi.\sum_{i\in A}e^{x_i}.

它已经减去局部最大值,从而避免大正数指数溢出。

因此 m 与 \ell 必须成对解释。

脱离 m 单独看 \ell,它没有统一的指数基准。

3. 两个局部最大值如何得到全局最大值

给两个不相交块 A、B:

mA=maxiAxi,mB=maxiBxi.m_A=\max_{i\in A}x_i, \qquad m_B=\max_{i\in B}x_i.

全局最大值为

m=max(mA,mB).m=\max(m_A,m_B).

这一步很直接。

字幕中的“局部最值再取最值”应准确理解为:

4. 分母为什么不能直接相加

局部分母分别是

A=iAeximA,\ell_A=\sum_{i\in A}e^{x_i-m_A},
B=iBeximB.\ell_B=\sum_{i\in B}e^{x_i-m_B}.

二者的基准分别是 mAm_AmBm_B

mAmBm_A\neq m_B,直接计算

A+B\ell_A+\ell_B

相当于把不同单位的量混在一起。

需要先改写到全局最大值 m 的基准:

exim=eximAemAm.e^{x_i-m} =e^{x_i-m_A}e^{m_A-m}.

因此

=emAmA+emBmB.\ell =e^{m_A-m}\ell_A +e^{m_B-m}\ell_B.
图 2

合并局部状态 (mb,b)(m_b,\ell_b) 时,先取全局最大值 m,再将每个局部分母乘 embme^{m_b-m} 后求和。

原视频 · 00:40 ↗

视频在 00:36 进入分母维护,在 00:43 给出局部分母乘指数缩放后求和的结构。

所有指数缩放因子都不大于 1:

mAm0,mBm0.m_A-m\le0, \qquad m_B-m\le0.

这保留了安全 Softmax 的数值稳定性。

5. 整体计算例子:[1,2,3,5]

视频在 00:54 取 logits

[1,2,3,5].[1,2,3,5].

全局最大值是 5。

安全指数为

[e4,e3,e2,e0].[e^{-4},e^{-3},e^{-2},e^0].

因此全局分母是

=e4+e3+e2+1.\ell =e^{-4}+e^{-3}+e^{-2}+1.
图 3

对 logits [1,2,3,5][1,2,3,5],安全 Softmax 统一减去全局最大值 5,分母为 e4+e3+e2+1e^{-4}+e^{-3}+e^{-2}+1

原视频 · 01:00 ↗

整行计算给出目标答案。

接下来要证明循环与规约都能恢复同一个 \ell

6. 顺序循环:先看 [1,2,3],再引入 5

对旧集合 A=[1,2,3]:

mA=3,m_A=3,
A=e2+e1+1.\ell_A=e^{-2}+e^{-1}+1.

新元素单独构成 B=[5]:

mB=5,m_B=5,
B=1.\ell_B=1.
图 4

单元素块 [5][5] 的局部状态是 mb=5m_b=5b=1\ell_b=1;与旧状态合并时,旧分母必须乘 emold5e^{m_{old}-5}

原视频 · 02:00 ↗

视频在 02:01 说明新元素 5 的局部状态,并在 02:08 开始用状态更新。

新的全局最大值:

m=max(3,5)=5.m=\max(3,5)=5.

分母:

=e35(e2+e1+1)+e551.\ell =e^{3-5}(e^{-2}+e^{-1}+1) +e^{5-5}\cdot1.

展开:

=e4+e3+e2+1.\ell =e^{-4}+e^{-3}+e^{-2}+1.

与整行计算完全相同。

7. 并行规约:合并 [1,2] 与 [3,5]

左块 A=[1,2]:

mA=2,m_A=2,
A=e1+1.\ell_A=e^{-1}+1.

右块 B=[3,5]:

mB=5,m_B=5,
B=e2+1.\ell_B=e^{-2}+1.

全局最大值仍是

m=5.m=5.

分母合并:

=e25(e1+1)+e55(e2+1).\ell =e^{2-5}(e^{-1}+1) +e^{5-5}(e^{-2}+1).

展开:

=e4+e3+e2+1.\ell =e^{-4}+e^{-3}+e^{-2}+1.
图 5

顺序循环把 [1,2,3][1,2,3] 与新元素 5 合并;并行规约把 [1,2][1,2][3,5][3,5] 合并,两条路径用同一状态算子得到相同结果。

原视频 · 01:40 ↗

视频在 01:31 对比单线程循环与两个线程的规约,在 01:42 指出两者结果相等。

8. 把合并写成一个算子

定义状态

zA=(mA,A),zB=(mB,B).z_A=(m_A,\ell_A), \qquad z_B=(m_B,\ell_B).

定义合并

zAzB=(m,emAmA+emBmB),z_A\oplus z_B =\left( m, e^{m_A-m}\ell_A+e^{m_B-m}\ell_B \right),

其中

m=max(mA,mB).m=\max(m_A,m_B).
图 6

任意分块只要各自保存局部最大值与对应安全分母,就能用相同的重标定公式合并;这正是并行规约成立的基础。

原视频 · 02:20 ↗

空集合可以用单位元表示:

z=(,0).z_\varnothing=(-\infty,0).

因为对任意有效状态 z,数学上有

zz=z.z\oplus z_\varnothing=z.

9. 为什么这个算子满足结合性

每个状态 (mA,A)(m_A,\ell_A) 都等价表示集合 A 的两个真实量:

mA=maxiAxi,m_A=\max_{i\in A}x_i,

以及

emAA=iAexi.e^{m_A}\ell_A=\sum_{i\in A}e^{x_i}.

合并只是把两边的指数和换到共同安全基准后相加。

集合并集本身满足结合律,因此精确实数数学下:

(zAzB)zC=zA(zBzC).(z_A\oplus z_B)\oplus z_C =z_A\oplus(z_B\oplus z_C).

这使 tree reduction、warp reduction 或分层 block reduction 成为可能。

10. 浮点下“相等”不一定逐位相同

实数数学下,循环与任意规约树得到完全相同的 m 与 \ell。 但浮点加法不满足严格结合律,(a+b)+c(a+b)+ca+(b+c)a+(b+c) 可能在最后几个 bit 上不同。

累加精度、块大小、归约树和编译器优化都会影响误差。

11. 与 FlashAttention 的衔接

在线 Softmax 只需维护每行的 m 与 \ell

FlashAttention 还要同时合并未归一化的输出分子:

O~A=iAeximAVi.\widetilde O_A =\sum_{i\in A}e^{x_i-m_A}V_i.

其重标定因子与分母完全一致:

O~=emAmO~A+emBmO~B.\widetilde O =e^{m_A-m}\widetilde O_A +e^{m_B-m}\widetilde O_B.

最终

O=O~.O=\frac{\widetilde O}{\ell}.

所以本课的 (m,)(m,\ell) 归约,是理解 FlashAttention 输出在线累加的直接前置。

跟练与练习

编者练习

用两种分组验证同一个分母 给定 logits [0,2,4].[0,2,4]. 分别按以下顺序合并: 证明最终状态相同。

  1. 先合并 [0,2],再合并 [4];
  2. 先合并 [0],再合并 [2,4]。
查看参考答案

第一种分组:
(mA,A)=(2,e2+1),(m_A,\ell_A)=(2,e^{-2}+1),
(mB,B)=(4,1).(m_B,\ell_B)=(4,1).
合并得
m=4,m=4,
=e24(e2+1)+1=e4+e2+1.\ell=e^{2-4}(e^{-2}+1)+1 =e^{-4}+e^{-2}+1.
第二种分组:
(mA,A)=(0,1),(m_A,\ell_A)=(0,1),
(mB,B)=(4,e2+1).(m_B,\ell_B)=(4,e^{-2}+1).
合并得
m=4,m=4,
=e041+(e2+1)=e4+e2+1.\ell=e^{0-4}\cdot1+(e^{-2}+1) =e^{-4}+e^{-2}+1.
两种分组都等于整行以最大值 4 为基准的安全分母。

常见误区

误区 1:局部分母可以直接相加

只有局部最大值相同,或都已经统一到同一全局最大值时才可以。

一般情况必须乘 embme^{m_b-m}

误区 2:运行最大值改变时,只更新新元素

错。

旧分母的所有贡献都要一起缩放到新基准。

误区 3:“动态规划”意味着要保存长度为 N 的 DP 表

本课只需常数个每行状态。

“动态规划”是局部状态递推的理解方式。

误区 4:规约与循环必须逐位相同

数学结果相同,浮点执行顺序不同可能产生微小舍入差异。

误区 5:最大值 m 本身就是 Softmax 分母

m 只是数值稳定性的指数基准。

真正分母是与 m 绑定的 \ell

误区 6:板书中的 sum 是普通 logits 求和

不是。

它表示安全指数之和:

=iexim.\ell=\sum_i e^{x_i-m}.

本课小结

  • 在线 Softmax 的每行状态是运行最大值 m 与安全分母 \ell
  • 两块合并时先取 m=max(mA,mB)m=\max(m_A,m_B)
  • 分母按 =emAmA+emBmB\ell=e^{m_A-m}\ell_A+e^{m_B-m}\ell_B 重标定。
  • 顺序循环是旧状态与新元素合并,规约是两个局部状态合并。
  • 精确数学下合并算子满足结合律,因此可并行规约。
  • 浮点下不同规约树可能不逐位相同,但应实现同一数学 Softmax。
  • 这一 (m,)(m,\ell) 状态合并是理解 FlashAttention 输出在线累加的基础。
12

单元综合

从分块等价到在线归约:FlashAttention 的前向与反向

单元能力目标

完成本单元后,应能把 FlashAttention 还原为一组可验证的数学等价变换与存储调度。

具体需要做到:

  • 证明按 QQ 行和 KK 行分块不会改变注意力矩阵的定义;
  • 从安全 Softmax 推导在线状态 (m,)(m,\ell)
  • 用可合并的归约解释顺序扫描、并行归约与分块计算;
  • 说明 FlashAttention 如何在片上存储分数块,并立即并入输出累加器;
  • 判断 causal mask 中哪些 tile 能整块跳过;
  • 区分 FlashAttention 版本调度变化与核心数学不变量;
  • 区分线性持久存储、平方级 dense 计算与实际 HBM 流量;
  • 用矩阵链式法则写出 dQ,dK,dVdQ,dK,dV
  • 解释反向为何能通过重计算避免持久化 N×NN\times N 中间梯度。

概念连接

1. 从完整注意力定义出发

对单个 attention head,先写出

S=QKdk,P=softmaxrow(S),O=PV.S=\frac{QK^\top}{\sqrt{d_k}}, \qquad P=\operatorname{softmax}_{row}(S), \qquad O=PV.

若有 causal mask MM,则实际分数为

S^=QKdk+M,\widehat S=\frac{QK^\top}{\sqrt{d_k}}+M,

其中未来位置对应 Mij=M_{ij}=-\infty

FlashAttention 没有改变这三个数学对象的定义,改变的是它们被计算、保存和丢弃的顺序。

2. 每个分数元素只依赖一对行向量

分数矩阵的元素是

Sij=qikjdk.S_{ij}=\frac{q_i^\top k_j}{\sqrt{d_k}}.

它只依赖 QQ 的第 ii 行和 KK 的第 jj 行。

所以,只要所有 (i,j)(i,j) 组合都被计算且不重不漏,就能恢复与整体矩阵乘法相同的 SS

3. 按 Q 行分块是纵向拼接

Q=[Q1Q2],Q= \begin{bmatrix} Q_1\\ Q_2 \end{bmatrix},

则矩阵乘法的分块法则直接给出

QK=[Q1KQ2K].QK^\top= \begin{bmatrix} Q_1K^\top\\ Q_2K^\top \end{bmatrix}.

每个 QbQ_b 只负责一组 query 行,但要遍历它们可见的全部 key,才能得到这些行的全局 Softmax 统计量。

4. 同时分 Q 块与 K 块

再将 KK 分为 K1,K2,K_1,K_2,\ldots,某个 tile 是

Sbc=QbKcdk.S_{bc}=\frac{Q_bK_c^\top}{\sqrt{d_k}}.

所有 (b,c)(b,c) tile 构成原分数矩阵的完整笛卡尔覆盖。

但此时一个 QbQ_b 对应的 Softmax 分母被分散在多个 KcK_c 中,需要在遍历 key tile 时逐块合并。

5. 安全 Softmax 避免指数溢出

对一行分数 x1,,xnx_1,\ldots,x_n,直接计算 exie^{x_i} 可能溢出。

安全写法为

m=maxixi,=iexim,pi=exim.m=\max_i x_i, \qquad \ell=\sum_i e^{x_i-m}, \qquad p_i=\frac{e^{x_i-m}}{\ell}.

传统的逻辑数组遍历是:

  1. 第一遍找 mm
  2. 第二遍计算 \ell
  3. 第三遍输出归一化概率。

这里的“三遍”是逻辑遍历,并不等同于某个具体 kernel 必然产生三次完整 HBM 读写。

6. 在线 Softmax 同时更新最大值和分母

扫描到新元素 xx 时,令

m=max(m,x).m'=\max(m,x).

原来的分母是围绕旧最大值 mm 表示的,需换到新基准 mm'

=emm+exm.\ell'=e^{m-m'}\ell+e^{x-m'}.

因此 (m,)(m,\ell) 就是一行 Softmax 的最小核心状态。

第一次扫描可以同时得到安全的 mm\ell,再扫描一次归一化,将三遍降为两遍。

7. 两段状态可以稳定合并

设两段分数的状态分别是 (mA,A)(m_A,\ell_A)(mB,B)(m_B,\ell_B)

合并时:

m=max(mA,mB),m=\max(m_A,m_B),
=emAmA+emBmB.\ell=e^{m_A-m}\ell_A+e^{m_B-m}\ell_B.

这正是将两段的安全指数和改写到同一个全局最大值下。

顺序在线循环可看作“已有状态”与“单个元素状态”反复合并;分块并行则是“块状态”之间的树形归约。

8. 数学可结合不代表浮点逐位一致

在精确实数算术中,上述合并对分段方式是结合的。

在浮点实现中,不同归约树会改变加法顺序,可能产生微小舍入差异。

所以需区分:

  • 数学上计算同一个 Softmax;
  • 数值上不承诺不同 kernel 逐 bit 一致。

9. FlashAttention 还需第三类状态

最终输出的一行是

o=jpjvj=jesjmvj.o=\sum_j p_jv_j =\frac{\sum_j e^{s_j-m}v_j}{\ell}.

定义未归一化分子累加器

o~=jesjmvj.\widetilde o=\sum_j e^{s_j-m}v_j.

当最大值从 moldm_{old} 变为 mnewm_{new} 时,\ello~\widetilde o 都必须乘同一重标定因子

emoldmnew.e^{m_{old}-m_{new}}.

于是 FlashAttention 每行维护三类概念状态:运行最大值、安全分母与输出分子。

10. 前向循环的不变量

固定一个 query tile QbQ_b 后,遍历所有可见的 (Kc,Vc)(K_c,V_c)

  1. Qb,Kc,VcQ_b,K_c,V_c 所需部分携入片上存储;
  2. 计算局部分数 SbcS_{bc}
  3. 应用 scale 和 mask;
  4. 计算局部最大值与指数和;
  5. 重标定旧状态,并入新块;
  6. 用局部权重与 VcV_c 更新输出累加器;
  7. 丢弃局部 SbcS_{bc},进入下一个 key tile。

遍历完成后再用 \ell 归一化。

关键不变量是:处理完前 cc 个 key tile 后,状态已精确表示这些 tile 联集上的 Softmax 与输出分子。

11. 局部分数矩阵不落入 HBM

普通实现若把整个 SSPP 写回 HBM,会产生 N×NN\times N 级别的持久存储和读写。

FlashAttention 让局部 SbcS_{bc} 只在 SRAM/寄存器等片上层级短暂存活,它与 VcV_c 的贡献并入输出状态后即可丢弃。

这是 IO-aware 调度:优化目标不是少算所有 dense 内积,而是避免反复物化并搬运巨大中间矩阵。

12. Causal mask 可以消除完全未来 tile

在 decoder causal attention 中,对于 query 位置 ii,key 位置 j>ij>i 必须被屏蔽。

对一个 (Qb,Kc)(Q_b,K_c) tile:

  • 若其所有 jj 都大于所有对应 ii,整块都是未来,可直接跳过;
  • 若 tile 与对角线相交,内部同时存在可见与不可见元素,仍需细粒度 mask;
  • 若 tile 完全在对角线下方,可整块正常计算。

跳过完全未来 tile 是精确消除零贡献,不是近似。

13. 单 token Decode 不等于普通 Prefill 的上三角

自回归 decode 中,新 query 位于当前最后,所有已缓存 key 都是它的过去或当前位置。

此时不存在同样的大面积未来区域。

多 token prefill 或块状验证则需在新块内保留 causal 约束。

因此“能跳过上三角”需结合 query/key 范围和实际 tile 边界判断。

14. FlashAttention-2 的外 Q 内 K/V 调度

观察到的高层循环是:

  1. 外层固定 QbQ_b 与其输出 ObO_b
  2. 内层流式遍历 Kc,VcK_c,V_c
  3. Qb,Ob,mb,bQ_b,O_b,m_b,\ell_b 尽量留在片上;
  4. 内层完成后一次写回 ObO_b

与高层上外 K/VK/VQQ 的安排相比,这能减少对 query、输出与在线状态的重复读写。

但它不表示 K/VK/V 只从 HBM 读一次,也不表示循环交换是 v2 的唯一优化。

实际 tile 大小、warp 分工、共享存储与并行策略还受 GPU 架构和 kernel 版本影响。

15. 平方分数矩阵是持久存储的压力源

Q,K,VRN×d,Q,K,V\in\mathbb R^{N\times d},

分数和概率矩阵都是

S,PRN×N.S,P\in\mathbb R^{N\times N}.

若持久化 SSPP,元素数为 N2N^2

Q,K,V,OQ,K,V,O 各自为 NdNd 元素。

对比一个 N×dN\times d 张量,比值为

N2Nd=Nd.\frac{N^2}{Nd}=\frac{N}{d}.

例如 N=100000,d=128N=100000,d=128时,比值约为 781.25781.25

16. 线性显存不意味线性计算

FlashAttention 的关键持久张量可随 NdNd 增长,并不需保存整个 N2N^2 分数矩阵。

但 dense attention 仍需计算大量 query-key 内积和权重-value 累加,主要算术量仍是

O(N2d).O(N^2d).

所以需同时报告:

  • 持久中间存储;
  • HBM 读写量;
  • 片上临时存储;
  • 计算量;
  • 硬件上的实际吞吐。

17. 反向先沿 O=PV 传播

设损失对输出的梯度为 dOdO

O=PV,O=PV,

矩阵微分是

dO=dPV+PdV.dO=dP\,V+P\,dV.

用 Frobenius 内积或 trace 把微分中的待求项移到末尾,得

PL=dOV,VL=PdO.\nabla_P L=dO\,V^\top, \qquad \nabla_V L=P^\top dO.

这里的乘法次序不能交换。

18. Softmax 反向是逐行 JVP

对某一行 p=softmax(s)p=\operatorname{softmax}(s),已知上游梯度 gpg_p,有

gs=p(gpgp,p1).g_s=p\odot\left(g_p-\langle g_p,p\rangle\mathbf 1\right).

即先求该行的标量

δ=jgp,jpj,\delta=\sum_j g_{p,j}p_j,

再计算

gs,j=pj(gp,jδ).g_{s,j}=p_j(g_{p,j}-\delta).

scale 1/dk1/\sqrt{d_k}、mask 与 dropout 都必须按前向的真实路径加入,不能在简化推导后遗漏。

19. 再沿 S=QK^T 传播

S=QK,S=QK^\top,

QL=dSK,KL=dSQ.\nabla_Q L=dS\,K, \qquad \nabla_K L=dS^\top Q.

若前向包含 1/dk1/\sqrt{d_k},相应梯度也需乘该系数。

用 shape 可以自检:

dS[N,N]K[N,d]dQ[N,d],dS[N,N]K[N,d]\rightarrow dQ[N,d],
dS[N,N]Q[N,d]dK[N,d].dS^\top[N,N]Q[N,d]\rightarrow dK[N,d].

20. 反向不必持久保存平方梯度

朴素理解可能会把 P,dP,dSP,dP,dS 都当作需要常驻的 N×NN\times N 张量。

FlashAttention 反向逐块重建局部中间量:

  1. Qb,KcQ_b,K_c 重算局部分数;
  2. 借助前向保存的行统计量重建局部 PbcP_{bc}
  3. 计算局部 dPbc,dSbcdP_{bc},dS_{bc}
  4. 立即累加到 dQb,dKc,dVcdQ_b,dK_c,dV_c
  5. 丢弃局部平方 tile。

最终梯度 dQ,dK,dVdQ,dK,dV 与输入一样为 N×dN\times d

线性额外存储是用重计算、分块调度和 IO 交换得到的,并不意味 backward 变成线性算术量。

对比与决策

1. 数学等价与逐位一致

  • 分块覆盖同一组 (i,j)(i,j) 内积,并用正确在线状态合并:数学目标相同。
  • 归约顺序和 kernel 不同:浮点结果可能有微小差异。
  • 漏 tile、重 tile、忘记重标定或 mask 错位:不再数学等价。

2. 显存复杂度与计算复杂度

  • 不持久化 N×NN\times N 中间量:可将关键额外存储降到线性量级。
  • 仍计算 dense query-key 组合:算术仍是平方级 token 耦合。
  • 运行更快:主要来自减少 HBM 物化/搬运和提高片上复用,不能仅由大 O 记号推断。

3. 哪些 causal tile 能跳过

  • 完全位于未来区:整块跳过。
  • 跨越对角线:计算 tile,内部细粒度 mask。
  • 完全位于可见区:正常计算。
  • 单 token decode:已有 key 通常全可见,不套用普通 prefill 的上三角图像。

4. 前向存储还是反向重计算

  • 保存更多 P/SP/S:反向重算少,但显存和 IO 大。
  • 保存紧凑行统计:反向逐块重构,显存少,增加重算与调度复杂度。

选择应根据显存瓶颈、算力余量、硬件带宽与 kernel 可用性,不是只比较某一个公式。

5. v1/v2 调度与不变的数学核心

高层循环顺序、并行分解和 warp 工作分配可以改变,但以下原则不变:

  • 完整覆盖所有可见分数;
  • 不同块的 Softmax 统计使用全局最大值重标定;
  • 输出分子与分母同步合并;
  • 局部平方 tile 不作为全局持久张量。

综合训练

编者练习 1

Q=[Q1;Q2]Q=[Q_1;Q_2]K=[K1;K2]K=[K_1;K_2]。写出 QKQK^\top2×22\times2 分块形式,并说明为什么单独计算 Q1K1Q_1K_1^\top 后就立即对其行做 Softmax 一般是错的。

查看参考答案

分块矩阵是
QK=[Q1K1Q1K2Q2K1Q2K2].QK^\top= \begin{bmatrix} Q_1K_1^\top & Q_1K_2^\top\\ Q_2K_1^\top & Q_2K_2^\top \end{bmatrix}.
Q1Q_1 的每一行需要与 K1K_1K2K_2 中所有可见 key 共同归一化。只对第一个 tile 单独 Softmax 会使其分母缺少 K2K_2 的指数项。正确做法是在两个 key tile 之间合并 (m,,o~)(m,\ell,\widetilde o)

编者练习 2

两个分数块的 Softmax 状态为 (mA,A)=(2,3),(mB,B)=(4,2).(m_A,\ell_A)=(2,3), \qquad (m_B,\ell_B)=(4,2). 求合并后的 (m,)(m,\ell)

查看参考答案

全局最大值是 m=4m=4
第一块的分母需从基准 2 换到基准 4:
=e243+e442=3e2+2.\ell=e^{2-4}\cdot3+e^{4-4}\cdot2 =3e^{-2}+2.
数值约为 2.4062.406。若有输出分子状态,第一块的 o~A\widetilde o_A 也需乘 e2e^{-2},第二块乘 11

编者练习 3

N=65536,d=128N=65536,d=128 的单头注意力,与一个 N×dN\times d 张量相比,N×NN\times N 分数矩阵的元素数多少倍?这个数字能否证明 FlashAttention 的计算复杂度是线性?

查看参考答案

比值为
N2Nd=65536128=512.\frac{N^2}{Nd}=\frac{65536}{128}=512.
这只比较两类张量的元素数,说明不物化分数矩阵的存储优势。Dense attention 仍需处理约 N2N^2 个 query-key 组合,主要算术量仍为 O(N2d)O(N^2d)

编者练习 4

已知 O=PV,S=QK,O=PV, \qquad S=QK^\top, 写出不经 Softmax 这一步时的 dP,dV,dQ,dKdP,dV,dQ,dK,并用 shape 检查结果。

查看参考答案

dORN×dvdO\in\mathbb R^{N\times d_v},则
dP=dOV,dV=PdO.dP=dO\,V^\top, \qquad dV=P^\top dO.
将逐行 Softmax JVP 得到的梯度记为 dSRN×NdS\in\mathbb R^{N\times N},则
dQ=dSK,dK=dSQ.dQ=dS\,K, \qquad dK=dS^\top Q.
右侧结果分别与 P,V,Q,KP,V,Q,K 的 shape 一致。若 SS 含 scale,还需传播 1/dk1/\sqrt{d_k}

进入下一单元前

  • 已能从 Sij=qikj/dkS_{ij}=q_i^\top k_j/\sqrt{d_k} 证明分块覆盖的完整性。
  • 已能独立推导在线 Softmax 的 (m,)(m,\ell) 更新与块合并公式。
  • 已能说明输出分子累加器为什么必须与分母同步重标定。
  • 已能判断 causal tile 的整块跳过与块内 mask 边界。
  • 已能区分 O(Nd)O(Nd) 持久存储与 O(N2d)O(N^2d) dense 计算。
  • 已能写出 attention 反向的矩阵梯度链条。
  • 已能解释 backward 通过逐块重计算避免持久化平方中间量。
  • 若仍把“分块”理解为近似,回看 P26、P27。
  • 若仍会在块内独立做 Softmax,回看 P28、P33、P36。
  • 若仍把线性显存写成线性计算,回看 P31、P32、P35。
  • 若对反向乘法次序不确定,回看 P34,并每次做 shape 自检。