返回关卡地图
预判 0 / 3 闯关题:0 / 3 XP 0 / 300

Level 20 | 零基础导学关卡

FlashAttention 模拟

FlashAttention Sim

先通过图解和小例子理解原理,再做预测与练习;不懂的地方可以反复尝试,最后把理解写进官方 Notebook。

本课官方 Notebook ↗
20_FlashAttention_Sim.ipynb 推理优化AttentionFlashAttention
Mission 1 概念练习

先建状态账本:分块改变计算顺序,不改变 Attention 答案

🎯 先猜一猜

seq_len=5、block_size=3,外层循环第二次执行 i=3。

q[3:6] 的实际 shape 是什么?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • 本题刻意省略 Batch 和 Head 维度,q、k、v 都是 [seq_len, dim];返回 out 也必须是 [seq_len, dim]。
  • 外层固定一块 Q,内层遍历所有 K/V 块。Python 切片允许终点超过长度,所以最后一块可以比 block_size 短。
  • 每个 query 行都维护自己的最大值 m 和指数和 l,因此它们是 [seq_len, 1] 列向量,便于向 score 的列方向广播。
  • q_block 已在外层乘 scale = 1/sqrt(dim),内层计算分数时不能再次缩放。

图解原理

把内层循环看成逐页读取 K/V:每读一页,就把这页对当前 Q 块的贡献合并进三个滚动状态。out 记当前归一化输出,m 记目前见过的最大 score,l 记以 m 为基准的指数和。翻完所有页后,这块 Q 的精确 Attention 才完成。

外层 i固定 q[i:i+block_size] 内层 j逐块读取 k 和 v 滚动合并更新 out_i、m_i、l_i 写回保存当前 Q 块的最终状态

标准 Attention 的大中间量

Q [S,D]Kᵀ [D,S]Scores [S,S]

完整分数矩阵会随序列长度按平方增长。

当前模拟只保留一个小分数块

Q_i [Bq,D]K_jᵀ [D,Bk]S_ij [Bq,Bk]

结果仍是精确 Attention,改变的是中间量的生存时间。

out[seq_len, dim];初始为 0,逐块形成最终输出m[seq_len, 1];初始为 -inf,表示尚未看过任何 scorel[seq_len, 1];初始为 0,表示尚无指数贡献scalePython float,值为 1 / sqrt(dim)device新张量至少跟随 q.device;可见测试使用 CPU float32
m 不能从 0 开始。若第一块 score 全是负数,0 会变成一个数据中从未出现的假最大值;-inf 才表示“空状态”。

语法热身:用分批处理成绩表练初始化和尾块切片

import torch

rows, width = 5, 3
scores = torch.randn(rows, width)

summary = torch.zeros((rows, width), device=scores.device)
running_max = torch.full((rows, 1), -float('inf'), device=scores.device)
running_sum = torch.zeros((rows, 1), device=scores.device)

for start in range(0, rows, 3):
    batch = scores[start:start + 3]
    print(batch.shape)  # [3,3],然后 [2,3]

从语法例子迁移到 TODO 1

  • summary:对应 out,行数与序列长度一致,列数与 value 特征维一致
  • running_max:对应 m;每一行只保存一个最大值,所以末维是 1
  • running_sum:对应 l;每一行只保存一个指数和,所以末维是 1
  • scores.device:对应 q.device,避免创建在错误设备上的状态张量

巩固一下

为什么 m 和 l 使用 [seq_len, 1],而不是 [seq_len]?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 2 概念练习

再合并一个分数块:最大值换基准,旧统计量必须重标定

🎯 先猜一猜

旧最大值 m_i=2,新块最大值 m_block=5,因此 m_new=5。

旧指数和 l_i 应乘哪个修正因子后再与新块相加?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • 矩阵乘法 [Bq,D] @ [D,Bk] 得到 [Bq,Bk];每行属于一个 query,每列属于当前 K 块中的一个 key。
  • Softmax 沿 key 维计算,所以 max 和 sum 都使用 dim=-1,并保留 keepdim=True。
  • m_new = maximum(m_i, m_block) 是逐元素比较两个 [Bq,1] 张量;它与 torch.max(S_ij, dim=...) 的归约作用不同。
  • P_ij = exp(S_ij - m_new) 已经把当前块放到新的共同基准下,因此 l_block 直接对 P_ij 求和。

图解原理

难点不是求新块的指数和,而是“最大值基准可能变了”。旧 l_i 是按旧 m_i 计算的;当 m_new 更大时,旧指数整体要乘 exp(m_i-m_new) 才能和新块相加。这就像两份用不同汇率记录的账,先换成同一种基准才能合并。

1S_ij = q_block @ k_block.T[Bq,D] @ [D,Bk] → [Bq,Bk]
2m_block = max(S_ij, dim=-1)当前 K 块每个 query 行的局部最大值
3m_new = maximum(m_i, m_block)选择旧块与新块共同的稳定基准
4P_ij = exp(S_ij - m_new)广播减法后再取指数
5l_new = l_i * exp(m_i-m_new) + sum(P_ij)修正旧账,再加入新块
变量shape检查重点
S_ij[Bq,Bk]q_block 已缩放,不重复乘 scale
m_block[Bq,1]dim=-1 且 keepdim=True
m_new[Bq,1]torch.maximum 做逐元素比较
P_ij[Bq,Bk]减 m_new 时按列广播
l_new[Bq,1]旧 l 先按新最大值修正
如果直接写 l_i + l_block,只在最大值从未变化时才碰巧正确。一旦新块出现更大 score,旧的指数和必须先乘修正因子。

语法热身:用两批传感器读数练在线 log-sum-exp

import torch

old_max = torch.tensor([[2.0]])
old_sum = torch.tensor([[1.5]])
new_values = torch.tensor([[1.0, 4.0, 3.0]])

batch_max = new_values.max(dim=-1, keepdim=True).values
joint_max = torch.maximum(old_max, batch_max)
batch_exp = torch.exp(new_values - joint_max)
batch_sum = batch_exp.sum(dim=-1, keepdim=True)
joint_sum = old_sum * torch.exp(old_max - joint_max) + batch_sum

从语法例子迁移到 Notebook TODO

  • new_values:对应 S_ij;每行是一位 query 对当前 K 块的分数
  • batch_max:对应 m_block,沿最后一维归约并保留列维
  • joint_max:对应 m_new,合并旧 m_i 与当前 m_block
  • batch_exp:对应 P_ij,所有分数都使用 m_new 这个共同基准
  • joint_sum:对应 l_new,先修正旧 l_i 再加 l_block

巩固一下

TODO 3 中 torch.max 与 torch.maximum 分别负责什么?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 3 概念练习

最后更新输出并写回:旧贡献和新贡献都要除以同一个 l_new

🎯 先猜一猜

当前 Q 块已经遍历完所有 K/V 块,但代码没有执行 out[i:i+block_size] = out_i。

函数最终返回的 out 会出现什么问题?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • out_i 是已经归一化的旧输出 [Bq,D],不能直接与 P_ij @ v_block 相加;两者当前所用的归一化尺度不同。
  • P_ij @ v_block 的 shape 是 [Bq,Bk] @ [Bk,D] = [Bq,D],与 out_i 同形状。
  • 每处理完一个 K/V 块,都要令 m_i=m_new、l_i=l_new,下一轮才会从最新状态继续。
  • 内层循环结束后,out_i、m_i、l_i 只覆盖当前 Q 切片;必须写回全局对应的 i:i+block_size 区域。

图解原理

输出更新可以拆成两部分:先把旧的归一化输出缩放到新分母下,再把当前块的加权 V 贡献除以同一个 l_new。这个“同分母再相加”是保持数学等价的关键。循环状态没更新或最终没写回,都会让后续块看见过期数据。

旧输出的权重

l_i * exp(m_i - m_new) / l_new

把以前已经归一化的 out_i 调整到新的最大值与总指数和下。

当前块的贡献

(P_ij @ v_block) / l_new

先按当前块指数权重聚合 V,再除以合并后的共同分母。

更新 out_i合并旧输出与当前 V 块 更新 m_i/l_i供下一个 K/V 块使用 结束内层当前 Q 块已看完全部 keys 写回全局保存 out、m、l 对应切片
case 1
S=8,D=4,block=2
case 2
S=5,D=3,block=3
case 3
S=3,D=2,block=1
误差
max abs diff < 1e-5

语法热身:用“旧平均 + 新一批数据”练同分母合并

import torch

old_mean = torch.tensor([[10.0, 20.0]])
old_weight = torch.tensor([[2.0]])
new_weighted_sum = torch.tensor([[9.0, 12.0]])
new_weight = torch.tensor([[3.0]])

total_weight = old_weight + new_weight
merged = old_mean * (old_weight / total_weight)
merged = merged + new_weighted_sum / total_weight
print(merged.shape)  # [1,2]

从语法例子迁移到 TODO 6

  • old_mean:对应 out_i,它已经是旧块归一化后的输出
  • old_weight:对应修正后的 l_i * exp(m_i-m_new)
  • new_weighted_sum:对应 P_ij @ v_block
  • total_weight:对应 l_new,旧、新贡献共享的分母
  • merged:对应更新后的 out_i,shape 仍为 [Bq,D]

巩固一下

为什么可见测试包含 seq_len=5、block_size=3?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

把理解变成自己的代码

准备好,去官方题目试一试

在本页用小例子建立直觉,再去官方 Notebook 完成实现。先运行你自己的测试,遇到困难时再查看官方提示与参考答案。

前往本课官方 Notebook ↗

这节课会遇到的代码对象

flash_attention_forward_sim · run_case

先找题目里的输入、输出与 TODO,再把本课的手算过程对应进去;以官方题目中的函数说明和测试为准。

练习来源:Datawhale 官方仓库 · 518cc45。这里的入口直接打开官方版本,不读取或分享你的本地 Notebook。

完成 3 道闯关题后,本关即算完成;作业 checklist 用来辅助你回 notebook 练习。
进度只保存在当前浏览器 localStorage,分享 HTML 不会带走你的记录。
上一关:L19 下一关:L21