公司真题库

【阿里巴巴】Attention 和 FFN 哪个更吃算力?不同序列长度下结论一样吗?

91学AI·2026/7/27·16 阅读

考察点

这道题出自阿里巴巴 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 本身。

评论 (0)

暂无评论,快来抢沙发吧!

91学AI

© 2026 91学AI · 按岗位学 AI 与大数据. All rights reserved.