FlashAttention 与 KV Cache:长上下文为什么这么吃显存
注意力的平方级显存怪兽 回到注意力的计算复杂度:序列里每个 token 要和所有其他 token 算注意力分数,这是 N² 的关系。当序列长度 N 是 1K,还好;但长上下文到 128K,N² 就从百万涨到一百六十亿——中间生成的注意力矩阵大到把显存撑爆。更麻烦的是,传统实现会把整个注意力矩阵物化在显存里,算完才扔掉,显存占用居高不下。这也是为什么长上下文训练和推理都贵得离谱。FlashAttention(2022 年由 Tri Dao 等提出)就是奔着这个瓶颈来的。 FlashAttention 怎么省显存 FlashAttention 的核心不是改变数学结果,而是改变计算的访存顺序。它把注意力分块(tiling):把长序列切成小块,在 GPU 的高速 SRAM 里逐块算、逐块累加,从不把完整的 N×N 矩阵写到慢的 HBM 显存里。因为数学上等价,输出结果和标准注意力一模一样,但显存占用从 O(N²) 降到约 O(N),同时因为减少了显存往返,速度还快了两三倍。这是典型的「IO 感知」优化——它赢在少搬数据,而不是少算。后来的 FlashAttention-2、3 继续把并行度和硬件适配做到极致,如今几乎所有大模型推理都默认启用它。 KV Cache:推理时的隐形显存大户 训练靠 FlashAttention 省显存,推理则要对付另一个怪兽:KV Cache。自回归生成时,为了不重复计算历史 token 的 Key 和 Value,会把它们缓存下来。问题是这个缓存随上下文长度 × 并发数 × 层数 × 头维度线性膨胀。一个 70B 模型、8 路并发、32K 上下文,光 KV Cache 就能吃掉二三十 GB——这就是很多服务「单聊流畅、一开并发就 OOM」的根源。vLLM 的 PagedAttention 正是用分页管理 KV Cache 来压碎片,前面文章讲过。 长上下文的真实成本账 理解了这两层,就能算清长上下文的账:上下文从 4K 拉到 128K,KV Cache 涨 32 倍,注意力计算也随长度增长。所以「支持百万上下文」的模型,实际部署成本远不止权重那点钱——它要求你为 KV Cache 预备巨量显存,并发还上不去。这也是为什么业界在搞 KV 量化、MQA/GQA(多查询注意力,让多个查询头共享 KV,直接把 KV Cache 砍到几分之一)这类优化。GQA 已成为现代大模型的标配,本质就是为了在不显著掉质量的前提下,把 KV Cache 这个隐形大户压下来。 ...