LogoCookLLM文档
LogoCookLLM文档
首页CookLLM

原理精讲

词元化
Tokenization 基础BPE 算法详解GPT 系列 TokenizerBPE 训练工程化
模型架构
Transformer LM
从 token ids 到 logitsEmbedding 与 LM Head
Attention 机制
Self-Attention 到 GQAAttention Sink
位置编码
位置编码基础RoPE 数学推导RoPE 代码实现长度外推
GPU 编程基础
GPU 架构基础张量布局Triton 入门:向量加法
FlashAttention
Flash Attention 原理详解从朴素实现到 Auto-TuningBlock Pointer 与多维支持Causal Masking 优化Grouped Query Attention反向传播实现
分布式训练
数据并行ZeRO 优化器全分片数据并行张量并行流水线并行多维混合并行
推理优化
KV CacheContinuous BatchingPagedAttention

动手训练

概述
预训练
预训练数据Tokenizer 训练模型架构数据流水线训练循环监控与验证
X (Twitter)
系统工程FlashAttention

Flash Attention 原理详解

会员专享

通过交互式可视化,深入理解 Flash Attention 的核心技术:内存瓶颈、Online Softmax、与分块矩阵乘法。

标准 Attention 的内存瓶颈

在深入 Flash Attention 的代码实现之前,我们必须先回答一个底层问题:为什么标准的注意力机制公式 Softmax(QKT)VSoftmax(QK^T)VSoftmax(QKT)V 在现代 GPU 上跑得还不够快?

登录以继续阅读

这是一篇付费内容,请登录您的账户以访问完整内容。

GPU 内存层级:SRAM 与 HBM

首先我们需要建立一个极其重要的概念:在 GPU 架构中,所有的计算(如矩阵加法、乘法)都必须在靠近核心的 SRAM(Shared Memory,共享内存) 中进行。

这意味着:哪怕你的显存(HBM)有 80GB 那么大,数据也必须先被“搬运”到几十 MB 大小的 SRAM 中,才能被计算核心处理。

标准实现的逻辑陷阱

在 PyTorch 等深度学习框架的“朴素”实现中,Attention 的计算过程被拆分成了多个独立的算子(Op)。这导致了一个严重的效率问题:

  1. 第一步(QKTQK^TQKT):GPU 把 QQQ 和 KKK 从 HBM 搬到 SRAM,计算出分数矩阵 SSS。
  2. 中间结果过大:由于 SSS 的形状是 (N,N)(N, N),对于长序列来说,这个矩阵大到 根本存不下。
# 标准 Attention 实现的 IO 噩梦
def standard_attention(Q, K, V):
    # 1. HBM -> SRAM(计算) -> HBM(存 S)
    S = Q @ K.T

    # 2. HBM(读 S) -> SRAM(计算) -> HBM(存 P)
    P = softmax(S)

    # 3. HBM(读 P) -> SRAM(计算) -> HBM(存 O)
    O =

这种 “搬进来 -> 算一下 -> 踢出去 -> 再搬回来” 的反复 I/O 往返,就是性能最大的杀手。

SRAM 与 HBM 的带宽差异

你可能会问,存回 HBM 再读回来会有多大影响?

速度差异对比

存储类型容量示例 (A100)带宽速度比喻
SRAM (共享内存)~20 MB~19 TB/sF1 赛车 🏎️
HBM (显存)40~80 GB~1.5 TB/s普通轿车 🚗

SRAM 的带宽通常比 HBM 高出约 10 倍以上。

瓶颈本质:IO 受限 (IO-bound)

由于这个速度鸿沟的存在,如果算法不断在 HBM 和 SRAM 之间搬运中间数据,就会出现一种尴尬的局面:

GPU 强大的计算核心大部分时间都在“空转”,苦苦等待数据从缓慢的 HBM 运送过来。

这种状态被称为 I/O 受限(IO-bound),即计算能力被内存传输速度拖了后腿。

量化感受一下:在 standard Attention 中,读写 N×NN \times NN×N 的中间矩阵 SSS 和 PPP 所花费的时间,远远超过了实际进行矩阵乘法计算的时间。

SRAM 的容量限制

既然 SRAM 这么快,为什么不把整个 Attention 矩阵都放在 SRAM 里算完再走?这里涉及到了物理与经济的刚性制约:

物理制约与成本

SRAM 的存储密度极低,导致成本极高。根据相关资料(如 FlashAttention 论文引用的背景):

  • 成本:制造一个 80GB 容量的 SRAM 存储器,成本可能高达 13,000 美元(估算数量级)。
  • 对比:同样容量的 HBM 仅需 2,000 美元。

容量极限

在实际硬件中,A100 的 HBM 可以达到 80GB,但 SRAM 通常只有 192 KB / SM(每个流式多处理器)。由于这种容量限制,你无法一次性把整个 N×NN \times NN×N 的注意力矩阵塞进 SRAM。当序列长度 NNN 增加时,中间矩阵的大小呈 O(N2)O(N^2)O(N2) 爆炸式增长。

核心思路:IO 复杂度优化

Flash Attention 的核心逻辑就在于:既然 SRAM 贵且小但快,HBM 大且便宜但慢,那么我们就必须放弃“全量读写”的幻想。

我们需要引入两个核心思想:

  1. 分块(Tiling):将数据切分成能塞进 SRAM 的“小方块”。
  2. 算子融合(Kernel Fusion):在 SRAM 内部,一气呵成完成 QKTQK^TQKT、Softmax 和 VVV 的乘法,只在最后一步将最终结果 OOO 写回 HBM。

避免中间矩阵落盘

通过这种方式,我们根本不生成(也不写入 HBM)那个巨大的 N×NN \times NN×N 中间矩阵。

方法HBM 读写量复杂度
Standard AttentionO(N2)O(N^2)O(N2)随着序列变长,IO 爆炸
Flash AttentionO(N)O(N)O(N)线性增长,极大节省带宽

Flash Attention

深入理解 Flash Attention 的原理与 Triton 实现

从朴素实现到 Auto-Tuning

编写第一个 Flash Attention Kernel,并利用 Auto-Tune 进行性能优化。

目录

标准 Attention 的内存瓶颈
GPU 内存层级:SRAM 与 HBM
标准实现的逻辑陷阱
SRAM 与 HBM 的带宽差异
速度差异对比
瓶颈本质:IO 受限 (IO-bound)
SRAM 的容量限制
物理制约与成本
容量极限
核心思路:IO 复杂度优化
避免中间矩阵落盘
Online Softmax 原理
离线算法的局限性
在线算法与动态修正
修正公式推导
数值演示:以序列 [3, 2, 5, 1] 为例
Flash Attention 的数学原理总结
分块矩阵乘法 (Tiling)
为什么要进行“分块”?
可视化演示:分块计算流程
观察重点
Tiling 与 Attention 的结合
循环策略对比:V1 vs V2
图解交互指南
分块 Attention 的 Softmax 修正
朴素实现:局部 Softmax 的局限
解决方案:在线重缩放 (Online Rescaling)
初始化
内循环: 遍历 K-Blocks
最终步骤: 归一化
完整算法伪代码
算法全貌
总结
(N,N)
SRAM
  • 被迫回写:GPU 只能被迫将这个巨型矩阵 SSS 从 SRAM “踢”出去,写回到慢速的 HBM 中。
  • 反复折腾:到了下一步计算 Softmax(S)Softmax(S)Softmax(S) 时,GPU 又得重新去 HBM 把刚才存进去的 SSS 再搬回 SRAM。
  • P
    @
    V
    return O