GPU物理架构与内存层级
问题:Tensor Core 适合什么计算模式,HBM 和 SRAM 的带宽差距有多大,数据为什么一旦反复搬运就会让算子很快变成 memory bound。
关键词: Tensor Core, SRAM, HBM
GPU架构代际演进(适配大模型的演进)
| 代际 | 关键引入 | 代表指标 |
|---|---|---|
| Volta V100 (2017) | 首创 Tensor Core,FP16 MMA | 开启深度学习混合精度时代 |
| Ampere A100 (2020) | TF32 /FP16/BF16、MIG、非对称稀疏化 | FP16 312 TFLOPS;HBM 1.5 TB/s;L2 40MB |
| Hopper H100 (2022) | 原生 FP8 、Transformer Engine、Thread Block Cluster + TMA(HBM→SRAM 异步直搬,不过寄存器) | FP8 1979 TFLOPS;HBM3 3.35 TB/s;NVLink 900 GB/s |
| Blackwell B200 (2024) | 第二代 Transformer Engine、原生 FP4、NVLink 5 代 | NVLink 双向 1.8 TB/s 级(平台实现有差异) |
本质驱动:Transformer 对混合精度矩阵计算 和极高显存带宽的持续攀升需求。
TensorCore vs Cuda Core
- CUDA Core(FP32/INT32) :每个时钟周期只能执行 1 个标量 FMA(Fused Multiply-Add):
d = a * b + c。 - Tensor Core(FP16/BF16/FP8) :专为矩阵乘法设计,单周期可执行一个完整的 4×4 矩阵 MMA(Matrix Multiply-Accumulate):
D = A * B + C。 - 精度策略:乘法用低精度(FP16/FP8)加速,累加用单精度(FP32)保证精度。
- 本质区别:不是"把标量 FMA 做快一点",而是把一批矩阵乘加打包成更大的 MMA 一次完成 → 在 GEMM(自注意力 + MLP 几乎全是)场景下算力碾压普通 CUDA Core。
内存层级金字塔(越靠近计算单元越快、容量越小)
- Registers(寄存器) :<1 周期,每线程几十个 32-bit;变量太多 → Register Spilling,数据回退到 Local Memory(物理位于 HBM)。
- Shared Memory (SRAM / 片上共享内存) :~19 TB/s(A100),每 SM 仅几百 KB;同 Block 内线程协作、交换数据的主通道。Triton 的重要作用:自动管理 SRAM 的分配与调度。
- L2 Cache:所有 SM 共享,几十 MB,缓冲 HBM 读写。
- HBM(全局显存) :40~80 GB,仅 1.5~3 TB/s;每次计算都走 HBM(如 PyTorch 原生多次小算子)→ 严重 Memory Bound。
判断依据:算术强度 Arithmetic Intensity = FLOPs / Bytes。强度低 = 每搬一次数据只做很少计算 → 更容易被 HBM 带宽卡住;强度高 → 计算单元更容易跑满。
FlashAttention 如何利用 SRAM 解决访存瓶颈
- 原生路径:S=QKᵀ 写回 HBM → 读回做 Softmax → 再写回 HBM → 读回与 V 乘 → O(N²) 反复读写 → OOM + 极慢。
- FlashAttention 四步 :
- Tiling(切块):把 Q, K, V 切成小块,刚好塞进几百 KB 的 SRAM;
- Fusion(SRAM 内完成一切):Q_block、K_block 进 SRAM,Tensor Core 算 S_block;
- Online Softmax(在线归约):SRAM 内直接更新局部最大值和指数和,不写回 S;
- 最后乘 V_block,最终结果写回 HBM。
- 结论 :HBM 读写从 O(N²) 压到接近 O(N) 。FlashAttention 不是减少计算量,而是通过 SRAM 缓解 Memory Bound。
PCIe vs NVLink(多卡节点内通信)
- PCIe(外围组件互连):传统插槽,PCIe Gen4 双向 64 GB/s;跨 GPU 通常要经过 PCIe Switch 甚至 CPU → 延迟高、带宽低。
- NVLink(NVIDIA 私有互连) :专为 GPU-to-GPU 设计。
- A100 NVLink 3.0:每条链路 50 GB/s,单卡 12 条,总双向 600 GB/s(≈PCIe 快近 10 倍);
- H100 NVLink 4.0:总双向 900 GB/s;
- Blackwell / NVLink 5:1.8 TB/s 级。
- NVSwitch:同机内 8 卡全互连(All-to-All)无阻塞通信 → 跑满 All-Reduce / All-Gather 极限带宽的硬件基础。
FlashAttention模拟
痛点:标准 Attention 为何撑不住长上下文
- 问题不只是矩阵乘法变大,更麻烦的是中间 attention score 矩阵 S = QKᵀ 按序列长度平方增长(O(N²))。
- 标准实现:整块 QKᵀ 写到显存 → 读回做 softmax → 再写回 → 读回加权求和。计算还没结束,显存与带宽已被中间结果拖住。
Step 1 · 标准 Softmax vs Online Softmax
标准 Softmax 三步:
- m = max(x)(防溢出);
- l = Σ e^(x−m);
- yᵢ = e^(xᵢ−m) / l。
→ 在算出所有 x 之前无法算出 m 和 l,必须先把所有 x(即整个 S=QKᵀ 矩阵)存下来 ------ 这是必须分块的根本原因。

Online Softmax :只看到部分数据也能持续更新局部最大值 m_new 和局部指数和 l_new;当新块最大值更大时,用一个数学技巧修正 之前算好的部分,无需重算前面的块。
更新公式(三个核心):
- m_new = max(m_old, m_block)
- l_new = l_old · e^(m_old − m_new) + l_block · e^(m_block − m_new)
(代码中等价实现:l_new = l_i * exp(m_i - m_new) + l_block,因为 P 已按 m_new 归一化) - O_new = O_old · (l_old · e^(m_old − m_new) / l_new) + e^(S_block − m_new) · V_block / l_new

Step 2 · 分块机制(Tiling)
- 在序列维度对 Q, K, V 分块:外层循环遍历 Q 块,内层循环遍历 K/V 块。
- 数学上完全等价的前提下,显存消耗从 O(N²) 降到 O(N)。
- 纯 PyTorch 用 for 循环模拟底层 C++ 内存块调度(也是理解 FA 的最佳方式)。

Step 3 · 代码实现框架(6 个 TODO ↔ 6 步)
python
# TODO 1: 初始化输出 O,全局最大值 m,全局指数和 l
# 提示: 先构造与 seq_len / dim 对齐的输出张量,再初始化 m 和 l
out = torch.zeros((seq_len, dim), device=q.device)
m = torch.full((seq_len, 1), -float('inf'), device=q.device)
l = torch.zeros((seq_len, 1), device=q.device)
python
# TODO 2: 计算当前块的未归一化分数 S_ij
S_ij = q_block @ k_block.transpose(-2, -1)
python
# TODO 3: 计算当前块的局部最大值 m_block,并求出新的全局最大值 m_new
m_block = torch.max(S_ij, dim=-1, keepdim=True)[0]
m_new = torch.maximum(m_i, m_block)
python
# TODO 4: 计算 P_ij = exp(S_ij - m_new)
P_ij = torch.exp(S_ij - m_new)
python
# TODO 5: 计算当前块的局部指数和 l_block,并更新全局指数和 l_new
l_block = torch.sum(P_ij, dim=-1, keepdim=True)
l_new = l_i * torch.exp(m_i - m_new) + l_block
python
# TODO 6: 更新输出 O_i(使用 Online Softmax 的修正公式)
out_i = out_i * (l_i * torch.exp(m_i - m_new) / l_new) + (P_ij @ v_block) / l_new
完整代码块
python
def flash_attention_forward_sim(q, k, v, block_size=2):
"""
纯 PyTorch 模拟 FlashAttention 前向传播。
假设没有 Batch 和 Head 维度,q, k, v 的形状都是 (seq_len, dim)。
"""
seq_len, dim = q.shape
# TODO 1: 初始化输出 O,全局最大值 m,全局指数和 l
# 提示: 先构造与 seq_len / dim 对齐的输出张量,再初始化 m 和 l
out = torch.zeros((seq_len, dim), device=q.device)
m = torch.full((seq_len, 1), -float('inf'), device=q.device)
l = torch.zeros((seq_len, 1), device=q.device)
scale = 1.0 / math.sqrt(dim)
# 外层循环:遍历 Q 的分块
for i in range(0, seq_len, block_size):
q_block = q[i:i+block_size] * scale
m_i = m[i:i+block_size]
l_i = l[i:i+block_size]
out_i = out[i:i+block_size]
# 内层循环:遍历 K, V 的分块
for j in range(0, seq_len, block_size):
k_block = k[j:j+block_size]
v_block = v[j:j+block_size]
# TODO 2: 计算当前块的未归一化分数 S_ij
S_ij = q_block @ k_block.transpose(-2, -1)
# TODO 3: 计算当前块的局部最大值 m_block,并求出新的全局最大值 m_new
m_block = torch.max(S_ij, dim=-1, keepdim=True)[0]
m_new = torch.maximum(m_i, m_block)
# TODO 4: 计算 P_ij = exp(S_ij - m_new)
P_ij = torch.exp(S_ij - m_new)
# TODO 5: 计算当前块的局部指数和 l_block,并更新全局指数和 l_new
l_block = torch.sum(P_ij, dim=-1, keepdim=True)
l_new = l_i * torch.exp(m_i - m_new) + l_block
# TODO 6: 更新输出 O_i(使用 Online Softmax 的修正公式)
out_i = out_i * (l_i * torch.exp(m_i - m_new) / l_new) + (P_ij @ v_block) / l_new
# 更新全局状态
m_i = m_new
l_i = l_new
# 写回全局变量
out[i:i+block_size] = out_i
m[i:i+block_size] = m_i
l[i:i+block_size] = l_i
return out
Step 4 · 工业演进 V1 → V2 → V3 → V4(面试高频)
- FlashAttention-1 (2022) · 打破显存墙:Tiling + Recomputation,空间复杂度 O(N²)→O(N)。局限:Thread Block 内 Non-Matmul 计算偏多;短 batch / 长序列下 Occupancy 不高。
- FlashAttention-2 (2023) · 算法级优化 + 多维并行:① 减少 Non-Matmul FLOPs(把更多算力留给 Tensor Core);② Sequence Parallelism(序列级并行),长文本推理 GPU 更易满载。
- FlashAttention-3 (2024) · 绑定 Hopper 的极限压榨 :① WGMMA 异步计算(Warp Group 级指令,Tensor Core 后台异步执行);② TMA (Tensor Memory Accelerator)硬件级搬运器(全局→共享内存,释放搬运线程);③ 2-Stage → Ping-Pong Pipeline 软件流水线掩盖访存延迟,计算与访存重叠。
- FlashAttention-4 · CuTeDSL 与 Blackwell 方向:更强调"代码生成 + kernel 组织"一体化优化,不再停留在数学公式改写。
FlashAttention的四个概念
- O(N²) vs O(N)------为什么长序列必挂;
- Memory Bound 与算术强度------瓶颈在"搬"不在"算";
- Tiling + SRAM------为什么分块就能省带宽(HBM 比 SRAM 慢 ~12×);
- Online Softmax 的"修正而非重算"------原理层面听懂即可。