跳到正文

ZeRO 与 FlashAttention 各优化什么?

区分注意力计算、KV Cache 与 DeepSpeed ZeRO 训练状态分片

原题:大模型中的注意力优化与训练状态分片分别解决什么瓶颈?请列举 FlashAttention、MQA/GQA、稀疏或线性注意力等计算优化方法,并解释 DeepSpeed ZeRO 为什么属于训练显存优化而不是注意力算法;比较它们在训练、prefill 和 decode 阶段的作用边界。

推理优化 · 字节真题

回答与解析

注意力瓶颈与优化路线

标准稠密 self-attention 的分数计算和矩阵规模随序列长度二次增长。优化需按目标分层:

  • 精确 kernel:FlashAttention/SDPA 用 tiling、online softmax 和融合减少 HBM I/O 与辅助存储,仍计算 exact dense attention。
  • 减少 KV heads:MQA/GQA 降低解码 KV Cache 和带宽,可能改变质量/并行权衡。
  • 限制连接:滑动窗口、块稀疏或全局+局部模式减少实际计算,但改变可见性,需要任务验证。
  • 近似/线性 attention:改变 attention 计算形式,以近似误差换复杂度收益。
  • 系统层:KV 量化、分页缓存、连续批处理、张量并行和算子融合优化服务效率。

FlashAttention 与 ZeRO 必须分开

FlashAttention 是 exact、IO-aware 的 attention kernel;它不把稠密 FLOPs 变成线性,HBM I/O 的界依 head dimension 与片上存储。ZeRO 则在数据并行训练中分片 optimizer states、gradients 和 parameters,降低每卡训练状态显存;它不改变注意力计算。ZeRO 的节省比例取决于 data-parallel world size、精度和 stage,不能固定写成 4x/8x。

训练选型同时看峰值显存和 tokens/s;推理还要拆分 prefill、decode、首 token 延迟、KV 容量与吞吐。在目标硬件和 shape 上基准测试,才能判断 kernel、稀疏结构或分片策略是否真正有效。

口语版讲法(约3分钟)

  • 按计算、缓存和系统层分类优化
  • 准确解释 FlashAttention
  • 把 ZeRO 放回训练状态分片
  • 按训练与推理阶段做 benchmark

回答注意力优化时,我会先按瓶颈分类。标准稠密 self-attention 的 QK 分数和显式矩阵随序列长度二次增长,所以第一类方法是在不改数学结果的情况下优化 kernel,例如 FlashAttention 和 fused SDPA。它们用分块、online softmax 和重计算减少 HBM 数据搬运与中间矩阵存储,计算的是 exact dense attention,不是近似注意力。主要 FLOPs 仍是二次量级,I/O 的理论界还与 head dimension 和片上存储有关。

第二类是减少解码缓存,例如 MQA 和 GQA 减少 KV heads,从而降低 KV Cache 容量和 decode 带宽。第三类是改变连接模式,像滑动窗口、块稀疏或全局加局部注意力,只计算一部分位置;这会改变模型能看到的信息,需要在任务上验证。第四类是线性或核近似 attention,用近似误差换复杂度。推理系统还可以做 KV 量化、分页缓存、连续批处理和算子融合。

ZeRO 必须单独讲。它不是注意力算法,而是数据并行训练的状态分片:不同 stage 分别分片 optimizer states、gradients 和 parameters,降低每张卡保存的训练状态。它不会改变 QK 相乘或 softmax 的复杂度。节省多少取决于 data-parallel world size、stage、精度、参数和激活占比,不能背成固定四倍或八倍。

训练时,FlashAttention 主要改善注意力中间存储和 I/O,ZeRO 主要改善模型训练状态;两者可以同时使用,但解决的是不同问题。推理时 ZeRO 通常不是核心表述,应重点看 prefill 的 attention、decode 的 KV 带宽和批处理调度。

最终我会在目标 GPU 上固定 dtype、batch、head dimension 和序列长度,分别测训练峰值显存与 tokens/s、prefill 延迟、decode 延迟、KV 容量和质量变化。只有这样才能判断应该用精确 kernel、换 GQA、采用稀疏结构,还是先修服务调度。

选择方案前我会先定位阶段。训练或 prefill 中,如果 attention 中间量和 HBM 流量占主导,优先尝试精确融合 kernel;decode 中若每步主要读取历史 KV,就看 GQA、KV 量化、分页和批调度;若目标任务确实需要超长上下文且二次计算不可接受,才评估滑窗、稀疏或近似结构带来的质量损失。

实验表也要把算法与系统变量拆开。固定模型权重比较 kernel,可以验证输出与梯度一致;更换 GQA 或稀疏连接则已经改变模型结构,需要重新训练或至少做质量评测。ZeRO 的对照应看训练状态显存和通信,不该拿它解释推理 attention 变快。明确每项优化改变了什么,组合方案时才知道收益能否叠加。

若 profiler 显示瓶颈在 tokenizer、网络传输或 CPU 调度,继续改 attention 不会改善端到端体验。因此我会同时报告 kernel 微基准和请求级指标,并说明优化覆盖了哪段链路。系统优化的目标是降低真实任务成本,而不是只让某个算子看起来更快。

关键一句:ZeRO 与 FlashAttention 可以叠加但优化对象不同,端到端瓶颈可能随序列长度和并行规模迁移。

核验来源

  1. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
  2. ZeRO: Memory Optimizations Toward Training Trillion Parameter Models
  3. Fast Transformer Decoding: One Write-Head is All You Need

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设训练一个长上下文模型时同时遇到模型状态显存不足和 attention 中间张量过大。你会怎样分别使用 ZeRO 与 FlashAttention,并用哪些指标确认两者各自解决了什么瓶颈?

  2. 问法 2 · 层层追问

    标准注意力的主要计算和存储瓶颈是什么?……FlashAttention 如何通过 tiling 与 online softmax 降低 HBM I/O?……ZeRO 分片的又是什么?……为什么 ZeRO 不属于注意力计算优化?

  3. 问法 3 · 直球技术

    列举注意力优化路线,并重点区分 FlashAttention 的精确 IO-aware kernel 与 ZeRO 的训练状态分片,说明二者在训练和推理中的适用边界。

同模块相关题目