Transformer 的哪个部分最占用显存?
Transformer 的哪个部分最占用显存?
这道题考的是 Transformer 训练与推理时显存占用的核心差异,属于原理分析题。面试官想看你能不能把显存瓶颈这件事讲清楚,是只看表面数字,还是真的理解背后原因。
我分四个板块来讲:显存花在哪 → 参数 vs 激活两个维度 → 训练 vs 推理差异 → 实际优化策略。
一、先破题——显存都花在哪儿了?
先搞清楚训练时显存由什么构成。
四部分:模型参数 + 梯度 + 优化器状态 + 中间激活。
参数和梯度好理解,就是模型本身的权重。优化器状态,比如 Adam 有两个 moment 矩阵,这部分也占不少显存。
中间激活是什么?
就是你数据过每一层时,计算出来的那些临时结果。矩阵乘法输出、Softmax 输出、LayerNorm 输出……这些不是权重,是计算过程中的"中间产物"。
类比一下,就像厨房备菜。食材(参数)是一次性买的,摆在冰箱里。但你切好的菜(中间激活)会摊满整个操作台,占的空间才是变动的、才是真正让你 OOM 的。
关键点:中间激活是显存中变动最大的部分。 参数大小是固定的,但中间激活随 batch size 和序列长度线性增长。这才是大模型训练时显存爆炸的真正原因。
二、参数层面 vs 中间激活层面——两个维度结论不同
这是最容易搞混的地方。
参数层面:FFN 占大头。
Transformer 大概 2/3 的参数都在 FFN(前馈网络)里。Attention 部分参数量相对少一些。所以从参数数量看,FFN 是大哥。
但从中间激活层面看:Self-attention 才是霸主。
具体数据大概是这样:
- Self-attention 激活:占 54%
- MLP 激活:占 38%
- LayerNorm 激活:占 8%
Self-attention 为什么这么吃激活?
因为 Attention 计算里有个 s×s 的矩阵(s 是序列长度)。每层都要存 Q、K、V 的激活,还要存 Attention Score 矩阵。这部分显存随序列长度平方增长。
类比买房。FFN 像建筑面积(参数),看着大但公摊少。Attention 像实际使用面积(中间激活),看着数字不大但每平米都很"实在",利用率高。
所以这道题有两个答案:
- 问参数谁最多 → FFN
- 问激活谁最多 → Self-attention
你得先确认面试官问的是哪个维度。
三、训练 vs 推理——显存瓶颈完全不同
这是面试经常被追问的点。
训练时,显存大头是中间激活。
训练需要反向传播。你前向传播时存下来的激活值,反向传播全要用。所以激活得一直留着,这部分显存占比能到 60-70%。
再加上梯度、优化器状态,显存直接起飞。
推理时,中间激活不需要存了,但有个新问题:KV Cache。
推理时你逐 token 生成。每个 token 都要计算 attention,但之前算过的 K、V 不用重算,直接缓存起来。
这时候显存占用变成了 2 × n_layers × seq_length × hidden_dim × batch_size。
长序列下,KV Cache 也会爆。
Self-attention 的 s×s 矩阵是核心瓶颈。不管训练还是推理,只要序列一长,这玩意儿显存就蹭蹭往上涨。
类比一下:
- 训练像直播,全程录像,每帧都要存
- 推理像点播,缓存关键帧,但关键帧本身也可能很大
四、面试加分——实际优化策略
光说不练不够,面试官喜欢追问"那你怎么解决"。
三个实用策略:
1. 减小 batch size
最直接的办法。激活随 batch 线性增长,batch 砍一半,激活显存少一半。
2. 长序列用 GQA / MQA 优化
标准 Multi-Head Attention 有 n_kv 个 Key/Value 头。GQA(Grouped Query Attention)和 MQA(Multi-Query Attention)把这个数量降下来,KV Cache 直接缩水。
MQA:所有 head 共用一组 K、V。显存省很多,但可能掉点性能。
GQA:n_kv 小于 n_q,但不为 1。折中方案。Llama 2 用的就是这个。
3. 混合精度训练
FP32 改成 FP16/BF16,激活精度降低,显存直接少一半。BF16 比 FP16 更稳定,大模型训练基本都用这个。
类比搬家。大箱子换成小箱子(batch),易碎品用泡沫仔细包好(精度),能装的更多。
面试怎么答
基础版(100-150字):
训练阶段,中间激活是显存最大消耗源,占比能达到 60-70%。单层来看,Self-attention 的中间激活约占 54%,MLP 占 38%,LayerNorm 占 8%。
需要注意的是,参数层面 FFN 占 2/3,但参数大小固定,而中间激活随 batch size 和序列长度线性增长,才是真正的可变显存瓶颈。
训练靠存激活,推理靠 KV Cache,长序列时 Attention 的 s×s 矩阵是核心瓶颈。
加分版(在基础版上加这些):
补充一点优化思路。减小 batch size 直接有效;长序列场景可以用 GQA/MQA 减少 KV Cache 显存;混合精度训练(FP16/BF16)能降低激活精度,节省将近一半显存。
这三个手段根据实际场景组合使用,能有效缓解 OOM 问题。
一句话总结
参数层面 FFN 最大,但训练时中间激活才是显存瓶颈,尤其是 Self-attention 模块的 s×s 矩阵;推理时 KV Cache 成为新的瓶颈,长序列优化绕不开 GQA/MQA 这类方案。
