| A100 | H100 | H200 | B200 | |
|---|---|---|---|---|
| SRAM 容量 | 20.25 MB | 33 MB | 33 MB | 37 MB |
| HBM 容量 | 40 GB | 80 GB | 141 GB | 192 GB |
| HBM 带宽 | 1.55 TB/s | 3.35 TB/s | 4.8 TB/s | 8.0 TB/s |
| 整块矩阵 Q、Kᵀ、softmax(QKᵀ/√d) | n×d > 32 MB n×n > 32 GB | 只能放 HBM |
|---|---|---|
| 单个 token 的 q、k、v、o 向量 | 4 × d ~ 几 KB | 塞得进 SRAM |
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 # 最后一次性除
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]
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]
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
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
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