考察点
这道题出自阿里巴巴 AI 基础设施开发暑期实习一面,看似概念题实则考量化分析:不能只说「长序列 attention 贵」,要能推出交叉点在哪、训练和推理工况是否一致。面试官想确认你有 FLOPs 估算的基本功,这对推理优化岗位是日常技能。追问会往 KV cache、FlashAttention、MoE 对结论的影响走。
参考答案
先把两部分的 FLOPs 账算出来
设序列长 n、隐藏维度 d、FFN 中间维度 4d(标准配置),只看矩阵乘法的浮点运算量(乘加各算一次,系数 2):
FFN:两个线性层 d→4d、4d→d,各 2·n·d·4d = 8nd²,合计 16nd²。如果是 SwiGLU(三个线性层),是 24nd²。
Attention 分两块:
- 线性投影:Q、K、V、O 四个投影各
2nd²,合计 8nd²(GQA 会略减 K/V 投影); - 注意力本身:QK^T 是
2n²d,softmax 后乘 V 又是2n²d,合计约 4n²d——这就是唯一的二次项。
所以一层总算力约 24nd² + 4n²d,FFN 的 16nd² 占线性项的大头。
交叉点在哪
attention 二次项超过 FFN 的条件是 4n²d > 16nd²,即 n > 4d。以 d=4096 的 7B 级模型为例,交叉点在 1.6 万 token 左右。
- 短序列(n 几千):FFN 是算力大头,参数量也占全模型约三分之二(attention 投影 4d² 对 FFN 8d²),这和大家直觉「FFN 是参数仓库」一致。
- 长序列(n 数万以上):attention 的二次项开始反超,n 到 128K 时 attention 计算量是 FFN 的数倍,成为绝对瓶颈——这就是长上下文训练贵、长文本推理慢的根源。
训练和推理的结论不完全一样
训练时上面的 FLOPs 分析基本成立(还要乘 3 倍算反向传播)。推理要分工况:
- Prefill(处理 prompt):和训练同构,长 prompt 下 attention 二次项是算力瓶颈。
- Decode(逐 token 生成):每步只算一个 token,矩阵乘退化成矩阵向量乘,此时瓶颈不是算力而是显存带宽——每步都要把全部权重和整个 KV cache 读一遍。FFN 权重占大头,决定每步的权重读取量;KV cache 随 n 线性增长,决定长序列时的增量开销。所以 decode 阶段的优化抓手是量化(少读字节)和 KV cache 压缩,而不是减 FLOPs。
这个分析在工程上怎么用
- FlashAttention 的价值要重新表述:它不改 FLOPs 总量,靠分块计算避免把 n×n 的注意力矩阵落显存,把访存量从 O(n²) 降到 O(n),长序列下提速数倍——因为实际瓶颈往往卡在访存而不是计算。
- GQA/MQA 砍的是 KV cache 的显存和带宽,对训练 FLOPs 几乎没影响,主要优化 decode。
- MoE 直接改变结论:FFN 激活参数按 top-k 缩小,FFN 的算力占比大幅下降,attention 的相对占比上升,长序列下瓶颈更早转移到 attention。
一句话收束
短序列看 FFN,长序列看 attention;训练看 FLOPs,decode 看带宽。脱离序列长度和工况谈「谁更吃算力」没有答案,这本身就是面试官想听到的第一句话。
可能的追问
1. 为什么长上下文训练成本是「超线性」的?
attention 的 n² 计算之外,activation 显存也随 n 增长,倒逼更激进的梯度checkpoint和并行切分(序列并行),通信开销叠加,实际成本增长比 n² 还陡。
2. 滑动窗口注意力、稀疏注意力改变了什么?
把二次项的有效 n 换成窗口 w,attention 变回近线性,长序列成本大降;代价是全局依赖弱化,需要层间信息传递或混合少量全局层补偿。
3. 线性注意力和状态空间模型(Mamba 类)呢?
它们把序列依赖的复杂度压到线性,长序列优势大;但精确检索类能力(从远处复制特定 token)弱于 softmax attention,所以主流方案是混合架构或继续优化 attention 本身。