注意力复杂度与长序列瓶颈:O(n²) 的重负

01-注意力与基础 进阶 约 20 分钟 #复杂度#O(n²)#长上下文#瓶颈 更新 2026-10-02
当前状态:未学
本文基于模型知识整理(生成时未联网核对),复杂度分析建议对照 FlashAttention 论文 §2 与长上下文综述复核。

一句话定义

自注意力的计算与内存都是 O(n²)(n×n 注意力矩阵):序列翻倍、成本×4——100 万 token 上下文的注意力矩阵有 10¹² 个元素,直接物化远超显存;这一平方瓶颈是"长上下文"成为独立研究方向(kp-021/025)的根本原因。

为什么重要

上下文长度是 LLM 的核心能力指标(读长文档、多轮对话、代码库理解)——而 O(n²) 直接给它封顶。FFAttention 等工程突破(kp-021)、滑动窗口等架构改造(kp-025)、Mamba 等线性架构挑战(kp-031)全部围绕这条瓶颈展开。理解瓶颈的精确位置(计算 vs 显存 vs 带宽),才能判断各类方案的真价值。

前置知识

kp-002(注意力计算流程);大 O 记号。

核心概念

  • 成本拆解(n 序列长、d 维度、L 层):

- 计算:QKᵀ 与 softmax·V 各 O(n²d)——注意力占大头(FFN 是 O(nd²),长序列时被反超); - 显存:朴素实现要物化 n×n 注意力矩阵——4k 序列(fp16)单层单头约 128MB 级、多层多头乘上去爆显存; - 带宽:GPU 是"计算快、搬数慢"(算术强度),注意力恰好是内存带宽受限型算子(FlashAttention 的突破口,kp-021)。

  • 长序列的数学天花板:128k 上下文的 n² 元素数≈1.7×10¹⁰——即使不物化(Flash),计算量本身也随平方增长;推理时 KV Cache 还额外线性增长显存(kp-023 的另一个瓶颈面)。
  • 瓶颈催生的三条路线(后续知识点的地图):

1. 不物化:FlashAttention(tiling+重算,精确但省显存,kp-021); 2. 稀疏/线性近似:滑窗、线性注意力(kp-021/025,近似的); 3. 换架构:状态空间模型/线性 RNN(kp-031,O(n) 的挑战者)。

原理与机制

为什么 n² 主要伤显存与带宽而非"算力":现代 GPU 每秒可做 10¹⁴ 次乘加,但 HBM 带宽只有 ~10¹² 字节/秒——n² 矩阵写回再读回的搬运才是真实瓶颈(计算利用率 <15%)。FlashAttention 的洞察:瓶颈不在 FLOPs 在 IO——分块计算让注意力矩阵只留在高速 SRAM 中,从"显存受限"变"计算受限"。

为什么简单降采样是伪解:截断/池化会丢长程信息(kp-008 分块的核心价值恰在精确对齐);合格方案必须保精度(Flash)或有理有据地近似(稀疏假设:多数 token 只需局部+少数全局锚点)。

训练与推理的不对称:训练并行算全序列(n² 一次付);自回归推理逐步生成——每步对新 token 的注意力随已生成长度增长,累计推理成本 O(n²) 且 KV 显存 O(n)(kp-023 的推理优化主题)。

图示

成本表 (n=序列长, d=维度):
  注意力: 计算 O(n²d) | 显存 O(n²) (朴素) | 带宽受限型
  FFN:    计算 O(nd²) | 长序列时被注意力反超
瓶颈三条出路:
  1. 不物化: FlashAttention (tiling+SRAM, 精确 kp-021)
  2. 稀疏/线性: 滑窗/线性注意力 (近似 kp-021/025)
  3. 换架构: SSM/Mamba O(n) (kp-031)
128k上下文: n²≈1.7×10¹⁰ 元素 → 直接物化不可行

直观类比

O(n²) 像"会议纪要记录每个人的发言与每个人的关系":10 人会议 100 条关系还好,10 万人会议就是 100 亿条——会议室(显存)先爆。FlashAttention 像"不让纪要落地、当场汇总"(SRAM 分块);稀疏注意力像"只记录与重点人物的关系";Mamba 像改用"滚动摘要"(每步只带一份状态总结)。

实例或案例

  • 上下文竞赛:GPT-3 的 2k → GPT-4 Turbo 的 128k → Gemini 的 1M——每一步背后都是 kp-021/025 技术的组合拳。
  • 代码助手:整仓代码理解需要 100k+ 级上下文——O(n²) 下的现实折衷是检索增强(只把相关片段放进窗口)。
  • 成本对比:同一模型 4k 与 128k 输入的 API 定价差 2~4 倍——平方瓶颈穿透到商业定价。

常见误区

  • 误区一:"FlashAttention 是近似算法"。它是精确注意力(数学结果相同),只是不物化中间矩阵(IO 感知重算);"近似"的是线性注意力/稀疏家族。
  • 误区二:"上下文长=能力同等覆盖全窗口"。有效上下文(模型真正利用的范围)常短于标称窗口(kp-025 的"lost in the middle"问题)。
  • 误区三:"只算 FLOPs 估算速度"。注意力是带宽受限——不看 IO 模型会严重高估 GPU 利用率(kp-021 的核心洞察)。

与其他知识点的关系

  • kp-002:被约束的计算本体。
  • kp-021/025/031:三条出路的技术正文。
  • kp-023:推理侧的 KV 显存瓶颈。

自测题

  1. 自注意力的计算、显存复杂度各是多少?瓶颈主要在 GPU 的哪个资源?

答:计算 O(n²d)、朴素显存 O(n²);主要瓶颈是显存带宽(注意力矩阵的搬运),非算力。

  1. FlashAttention 解决的是哪一层瓶颈?代价是什么?

答:IO 瓶颈(不物化 n² 矩阵、SRAM 分块+重算)——代价是部分计算重复(重算策略)与实现复杂度,结果精确。

  1. 为什么长文档任务常用"检索增强+短窗口"而非超长窗口?

答:n² 成本与有效上下文问题(中间信息利用差)——检索把"平方"换成了"相关片段的短窗口"。

延伸阅读

  • Dao 等, "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness"(2022)。
  • 《Lost in the Middle》论文(长上下文利用率问题)。
  • 《Efficient Transformers: A Survey》(Tay 等,变体全景图)。