【llm-algo-leetcode学习笔记】显存与性能认知底座

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。
内存层级金字塔(越靠近计算单元越快、容量越小)
  1. Registers(寄存器) :<1 周期,每线程几十个 32-bit;变量太多 → Register Spilling,数据回退到 Local Memory(物理位于 HBM)。
  2. Shared Memory (SRAM / 片上共享内存) :~19 TB/s(A100),每 SM 仅几百 KB;同 Block 内线程协作、交换数据的主通道。Triton 的重要作用:自动管理 SRAM 的分配与调度
  3. L2 Cache:所有 SM 共享,几十 MB,缓冲 HBM 读写。
  4. 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 四步
    1. Tiling(切块):把 Q, K, V 切成小块,刚好塞进几百 KB 的 SRAM;
    2. Fusion(SRAM 内完成一切):Q_block、K_block 进 SRAM,Tensor Core 算 S_block;
    3. Online Softmax(在线归约):SRAM 内直接更新局部最大值和指数和,不写回 S;
    4. 最后乘 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 三步

  1. m = max(x)(防溢出);
  2. l = Σ e^(x−m);
  3. 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 的"修正而非重算"------原理层面听懂即可。

参考学习链接

GPU Architecture and Memory

FlashAttention Sim

相关推荐
rannn_1111 小时前
【力扣hot100】二叉树专题+总结
java·算法·leetcode·二叉树
啊嘞嘞?1 小时前
力扣(岛屿数量)
算法·leetcode
zander2582 小时前
LeetCode 198. 打家劫舍
算法·leetcode·深度优先
for_ever_love__2 小时前
python基础语法学习: 类型注解
开发语言·python·学习
摇滚侠2 小时前
《Docker技术入门与实战 第4版》阅读笔记 6 使用 Dockerfile 创建镜像 2
java·笔记·docker
吃着火锅x唱着歌2 小时前
LeetCode 648.单词替换
算法·leetcode·职场和发展
树的枝2 小时前
【STM32】05.TIM定时器
笔记·stm32·单片机·嵌入式硬件
Elsa️7462 小时前
leetcode 14.最长公共前缀
算法·leetcode·职场和发展
不会代码的小猴2 小时前
6. Qt网络编程
开发语言·c++·笔记·qt·算法