FlashAttention:为什么加速推理的同时还能节省显存
GPU 显存是分层的:层级越高越快,也越小
GPU SRAM GPU HBM Main Memory (CPU DRAM) SRAM :19 TB/s(20 MB) HBM :1.5 TB/s(40 GB) DRAM :12.8 GB/s(>1 TB)
A100 H100 H200 B200
SRAM 容量 20.25 MB33 MB33 MB37 MB
HBM 容量 40 GB80 GB141 GB192 GB
HBM 带宽 1.55 TB/s3.35 TB/s4.8 TB/s8.0 TB/s
想让 GPU 跑得快,就得尽量在 SRAM 里算完,少在中途从 HBM 读写数据。
假设提示词有 n 个 token、每个 token 向量 d 维 —— Q、K、V 各是一个 n×d 矩阵
\[\text{Attention}(\textcolor{#7c5cf0}{Q},\textcolor{#0e8fd4}{K},\textcolor{#e2691b}{V})=\text{softmax}\!\left(\frac{\textcolor{#7c5cf0}{Q}\textcolor{#0e8fd4}{K^{\top}}}{\sqrt{d_k}}\right)\textcolor{#e2691b}{V}\]
Q
n × d
×
K⊤
d × n
=
softmax(QK⊤/√d)
n × n
×
V
n × d
=
O
n × d
n 一般在 100000 量级还在往上涨,d 一般是 32 / 64 / 128
整块矩阵 Q、Kᵀ、softmax(QKᵀ/√d) n×d > 32 MB n×n > 32 GB 只能放 HBM
单个 token 的 q、k、v、o 向量 4 × d ~ 几 KB 塞得进 SRAM
深入矩阵计算内部:注意力分数绝对值 → Softmax 相对值 → 加权 v 向量
1\(S=QK^{\top}\)
第一行的 n 个值,就是第一个 token 对其他所有 token 的注意力绝对值分数
\(k_1\)
\(k_2\)
⋯
\(k_N\)
\(q_1\)
\(q_2\)
⋮
\(q_N\)
\(q_1k_1\)
\(q_1k_2\)
⋯
\(q_1k_N\)
\(q_2k_1\)
\(q_2k_2\)
⋯
\(q_2k_N\)
⋮ ⋮ ⋱ ⋮
\(q_Nk_1\)
\(q_Nk_2\)
⋯
\(q_Nk_N\)
2\(P=\text{Softmax}(S)\)
把注意力绝对值分数,归一化为总和为 1 的概率分布分数
\(s_1 =\)
\(q_1k_1\)
\(q_1k_2\)
⋯
\(q_1k_N\)
↓ ↓ ↓ \(p_1 =\)
\(\frac{e^{q_1k_1}}{\sum_x e^{q_1k_x}}\)
\(\frac{e^{q_1k_2}}{\sum_x e^{q_1k_x}}\)
⋯
\(\frac{e^{q_1k_N}}{\sum_x e^{q_1k_x}}\)
3\(O=PV\)
输出向量等于所有 token 对应的 Value 向量用注意力分布系数加权后累加
\(p_1\)
\(\times\)
\(v_1\)
\(v_2\)
⋮
\(v_N\)
\(=\)
\(v_1\)×\(p_{11}\) + \(v_2\)×\(p_{12}\) + ⋯ + \(v_N\)×\(p_{1N}\)
\(=\)
\(o_1\)
注意力矩阵计算的瓶颈:大量 GPU 消耗在等待 HBM
Memory: HBMCompute Memory: HBMCompute 每步矩阵计算都把结果写回 HBM, 然后下一步再从 HBM 读出来 能不能只从 HBM 读一次、 全部计算完再把最终结果写回 算子融合(Kernel Fusion)
一个理论上的融合算子 —— 但 \(e^{x}\) 会溢出
fused_attention.py
def fused_attention(q, K, V):
    d = 0                        # softmax 的分母
    o = zeros_like(V[0])         # 输出向量
    scale = 1 / sqrt(K.shape[1])  # 公式里的 1/√d_k
    for k, v in zip(K, V):       # 不断从HBM里读取k和v
        x = q @ k * scale
        o += exp(x) * v          # 分子直接累加进输出
        d += exp(x)              # 分母同步累加
    return o / d                 # 最后一次性除
\(\text{softmax}(x_i)\) \(=\) \(\dfrac{e^{x_i}}{\sum_{j} e^{x_j}}\) 实际计算时,\(x\) 太大会导致 \(e^{x}\) 数值溢出 \(\text{safe softmax}(x_i)\) \(=\) \(\dfrac{\textcolor{#10a37f}{e^{-m}}\cdot e^{x_i}}{\textcolor{#10a37f}{e^{-m}}\cdot\sum_{j} e^{x_j}} \;=\;\dfrac{e^{x_i\textcolor{#10a37f}{-m}}}{\sum_{j} e^{x_j\textcolor{#10a37f}{-m}}}\)
取 \(m=\max(x_i)\),每个 \(x_i-m\le 0\),\(e^{x_i-m}\) 全落在 \((0,1]\),指数再也不会溢出
算子融合的关键:在线Softmax
Safe Softmax
safe_softmax_3pass.py
def safe_softmax(x):
    m, d = -inf, 0
    # pass 1
    for xi in x:
        m = max(m, xi)
    # pass 2
    for xi in x:
        d += exp(xi - m)
    # pass 3
    return [exp(xi - m) / d for xi in x]
\(d_i = \sum_{j\le i} e^{x_j-m_N}\)
每一项都减总体最大值 \(m_N\),要整行扫完才知道
在线Softmax
online_softmax_2pass.py
def online_safe_softmax(x):
    m, d = -inf, 0
    # pass 1
    for xi in x:
        m_old = m
        m = max(m, xi)
        d = d * exp(m_old - m) + exp(xi - m)
    # pass 2
    return [exp(xi - m) / d for xi in x]
\(d'_i = \sum_{j\le i} e^{x_j-\textcolor{#10a37f}{m_i}}\)
每一项都减前 \(i\) 项的最大值 \(m_i\)
\(d'_i\) 和 \(d'_{i-1}\) 的递归式:
\(d'_i = d'_{i-1}\,\textcolor{#0a7d61}{e^{m_{i-1}-m_i}} + e^{x_i-m_i}\)
点击展开推导
把 \(d'_{i-1}\) 的定义代回去:
\[\begin{aligned} &= \Big(\sum\nolimits_{j\le i-1} e^{x_j-m_{i-1}}\Big)\textcolor{#0a7d61}{e^{m_{i-1}-m_i}} + e^{x_i-m_i}\\[4pt] &= \sum\nolimits_{j\le i-1} e^{x_j-m_i} + e^{x_i-m_i}\\[4pt] &= \sum\nolimits_{j\le i} e^{x_j-m_i} \end{aligned}\]
FlashAttention:继续用在线Softmax的思路建立递归式
Safe Softmax Attention
attention_2pass.py
def attention(q, K, V):
    m, d = -inf, 0
    o = zeros_like(V[0])
    scale = 1 / sqrt(K.shape[1])
    # pass 1
    for k in K:
        x = q @ k * scale
        m_old = m
        m = max(m, x)
        d = d * exp(m_old - m) + exp(x - m)
    # pass 2
    for k, v in zip(K, V):
        o += exp(q @ k * scale - m) / d * v
    return o
\(\boldsymbol{o}_i = \sum_{j\le i}\dfrac{e^{x_j-m_N}}{d_N}\,\boldsymbol{v}_j\)
\(m_N\)、\(d_N\) 要整行扫完才知道
FlashAttention
flash_attention_1pass.py
def flash_attention(q, K, V):
    m, d = -inf, 0
    o = zeros_like(V[0])
    scale = 1 / sqrt(K.shape[1])
    # pass 1
    for k, v in zip(K, V):
        x = q @ k * scale
        m_old = m
        d_old = d
        m = max(m, x)
        d = d_old * exp(m_old - m) + exp(x - m)
        rescale = d_old * exp(m_old - m) / d
        o = o * rescale + exp(x - m) / d * v
    return o
\(\boldsymbol{o}'_i = \sum_{j\le i}\dfrac{e^{x_j-\textcolor{#10a37f}{m_i}}}{\textcolor{#10a37f}{d'_i}}\,\boldsymbol{v}_j\)
每一项都用前 \(i\) 项的 \(m_i\)、\(d'_i\)
\(\boldsymbol{o}'_i\) 和 \(\boldsymbol{o}'_{i-1}\) 的递归式:
\(\boldsymbol{o}'_i = \boldsymbol{o}'_{i-1}\,\textcolor{#0a7d61}{\dfrac{d'_{i-1}\,e^{m_{i-1}-m_i}}{d'_i}} + \dfrac{e^{x_i-m_i}}{d'_i}\,\boldsymbol{v}_i\)
点击展开推导
把 \(\boldsymbol{o}'_{i-1}\) 和 \(d'_{i-1}\) 的定义代回去:
\[\begin{aligned} &= \Big(\sum\nolimits_{j\le i-1}\frac{e^{x_j-m_{i-1}}}{d'_{i-1}}\boldsymbol{v}_j\Big)\textcolor{#0a7d61}{\frac{d'_{i-1}\,e^{m_{i-1}-m_i}}{d'_i}} + \frac{e^{x_i-m_i}}{d'_i}\boldsymbol{v}_i\\[4pt] &= \sum\nolimits_{j\le i-1}\frac{e^{x_j-m_i}}{d'_i}\boldsymbol{v}_j + \frac{e^{x_i-m_i}}{d'_i}\boldsymbol{v}_i\\[4pt] &= \sum\nolimits_{j\le i}\frac{e^{x_j-m_i}}{d'_i}\boldsymbol{v}_j \end{aligned}\]
Tiling:将QKV分成一段段的小矩阵批量处理
QK⊤: N × N K⊤: d × N Q: N × d V: N × d sm(QK⊤)V 拷贝到 SRAM 拷贝 拷贝 在 SRAM 上算这一块 写回 HBM 内层循环 外层循环 外层循环 内层循环
绿 = 待在 HBM 的整块,橙 = 当前搬进 SRAM 的小块
flash_attention_tiled.py
def flash_attention_tiled(Q, K, V, tile_size):
    O = zeros_like(Q)
    scale = 1 / sqrt(K.shape[1])
    # 外层:搬 tile_size 行 Q 进 SRAM
    for i in range(0, len(Q), tile_size):
        Q_tile = Q[i:i+tile_size]
        # 每块 query 各持一份状态
        m = full((len(Q_tile), 1), -inf)
        d = zeros((len(Q_tile), 1))
        o = zeros_like(Q_tile)
        # 内层:扫 tile_size 行 K、V
        for j in range(0, len(K), tile_size):
            K_tile, V_tile = K[j:j+tile_size], V[j:j+tile_size]
            # tile_size × tile_size 的一小块分数
            x = Q_tile @ K_tile.T * scale
            m_old = m
            d_old = d
            m = maximum(m, x.max(1, keepdims=True))
            p = exp(x - m)
            d = d_old * exp(m_old - m) + p.sum(1, keepdims=True)
            rescale = d_old * exp(m_old - m) / d
            o = o * rescale + (p / d) @ V_tile
        # 一块算完才写回 HBM
        O[i:i+tile_size] = o
    return O