训练显存快速估算 (1)
训练显存快速估算¶
1. 训练一个7B参数的模型,FP16混合精度+AdamW,大概需要多少显存?¶
训练一个 7B(70 亿)参数的大语言模型,使用 FP16 混合精度和 AdamW 优化器,在未使用任何内存节省技术(如 LoRA、量化、ZeRO)之前,每个参数在训练时大约需要 20 字节的显存。这个数字是模型权重、梯度、以及优化器状态的总和。
具体来说,对于 7B 模型,每个参数 2 字节用于存储模型权重(FP16),2 字节用于存储梯度(FP16),另外,AdamW 优化器需要为每个参数存储一个 FP32 的主权重副本(4 字节),以及一阶动量(m)和二阶动量(v),各占 4 字节。因此,优化器状态总计为 $ 4 + 4 + 4 = 12 $ 字节。这样,权重、梯度和优化器状态加起来,每个参数需要 $ 2 + 2 + 12 = 16 $ 字节。所以,7B 模型这三项的总显存开销就是 $ 7B \times 16 $ 字节 = 112 GB。
然而,总显存需求还要加上前向传播过程中产生的“激活值”(Activations)。激活值的大小取决于批量大小(batch size)、序列长度(sequence length)和模型架构。在典型的微调设置下(例如,批量大小设为1,序列长度为2048),激活值通常会在上述112 GB的基础上再增加约20%到40%。因此,训练一个7B模型的总显存需求大约在130 GB到160 GB之间。
这就是为什么社区中流传着一个经验公式:“参数量 × 20”。7B 乘以 20 正好是 140 GB,这通常足以覆盖权重、梯度、优化器状态和一个典型规模的激活值开销。

模型参数占用
7B 参数 $ \times $ 2字节(FP16)
= 14 GB
AdamW优化器状态
每个参数2个状态×2字节
7B×4字节=28 GB
激活弹度&临时张量
≈ 14 GB
(10~15GB区间)
总息存估算
14 + 28 + 14 = 56 GB
实际建议:至少60 GB 显存

- 上述估算的 20 倍参数量的经验公式是如何推导出来的?列出各部分占比。
这个经验公式的推导基于混合精度训练中,模型对显存的典型消耗模式。为了清晰地展示,我们可以通过对一个1B(10亿)参数模型的分解来理解。
| 显存消耗部分 | 存储精度 | 每参数字节数 | 1B 参数显存 (GB) | 占比 |
| 模型权重 | FP16 | 2 Bytes | 2 GB | ~10% |
| 梯度 | FP16 | 2 Bytes | 2 GB | ~10% |
| 优化器状态 | FP32 | 12 Bytes | 12 GB | ~60% |
| 主权重副本 | FP32 | 4 Bytes | 4 GB | ~20% |
| 一阶动量 (m) | FP32 | 4 Bytes | 4 GB | ~20% |
| 二阶动量 (v) | FP32 | 4 Bytes | 4 GB | ~20% |
| 基础开销合计 | - | 16 Bytes | 16 GB | ~80% |
| 激活值 (估算) | FP16 | (经验值) | 4 GB | ~20% |
| 总计 | - | ~20 Bytes | 20 GB |
推导过程:¶
-
基础开销(权重+梯度+优化器):如上表所示,这部分的“硬成本”是固定的,对于1B参数模型是16 GB。优化器状态占比高达60%,远超权重本身。
-
激活值估算:激活值没有固定的“每参数”占用,因为它取决于 batch size 和序列长度。但在标准的微调设置下,经验表明激活值的大小大约是模型“基础开销”的 25% 左右。1B 模型的基础开销是 16 GB,其 25% 大约是 4 GB。
将基础开销(16 GB)和激活值估算(4 GB)相加,正好得到20 GB。因此,得出经验法则:每10亿参数大约需要20 GB显存。
对于 7B 模型, $ 7 \times 20\,GB = 140\,GB $。这个公式提供了一个快速、实用的上线估算,但精算时仍需根据具体的 batch size、序列长度以及是否使用 FlashAttention 等进行调整。
3. 对于 13B 模型,用 8 张 A100 80GB,能否训练?如果不能,显存瓶颈在哪?¶
理论上,8 张 A100 80GB 完全足以训练一个 13B 参数的模型,但前提是必须使用 DeepSpeed ZeRO 或模型并行等显存优化策略。如果采用最基础的分布式数据并行(DDP),则根本无法运行。
我们来分析不同策略下的情况:
普通数据并行(DDP,不可行)¶
在DDP模式下,集群中的每张GPU都要维护一份完整的模型副本,包括权重、梯度和优化器状态。
对于 13B 模型,这三项基础开销为 $ 13 \times 16\text{ Bytes} = 208\text{ GB} $。这已经远超单张 A100 的 80 GB 物理显存限制,连模型都无法加载,更不用说还要为激活值预留空间。
ZeRO Stage 1(可行,但非常拮据)¶
ZeRO Stage 1(可行,但非常拮据)
ZeRO-1 仅将优化器状态在所有 GPU 之间进行分片。
对于 13B 模型,其优化器状态大小为 $ 13 \times 12 $ Bytes = 156 GB。平均分到 8 张 GPU 上,每张卡持有 $ 156 / 8 = 19.5 $ GB。
但是,模型权重(26 GB)和梯度(26 GB)仍然需要每张卡完整保存。
这样一来,每张 GPU 的基础显存占用就是 26 (权重) + 26 (梯度) + 19.5 (分片后的优化器状态) = 71.5 GB。
此时,单卡剩余的可用显存仅有 80 - 71.5 = 8.5 GB。这点空间用于存放训练时的激活值是极度紧张的,可能只能支持极小的 batch size(如 1 或 2)或极短的序列长度,稍有波动就会 Out Of Memory (OOM)。因此,ZeRO-1 虽然理论上能跑,但实际应用中风险很高,不太实用。
ZeRO Stage 2(可行且推荐)¶
ZeRO-2 在 ZeRO-1 的基础上,进一步将梯度也进行了分片。
这意味着每张 GPU 现在只需持有 26 GB(权重)+ (26 / 8 ≈ 3.25 GB)(分片梯度)+ 19.5 GB(分片优化器状态)= 48.75 GB。
这样,单卡就空余出了 80 - 48.75 = 31.25 GB 的显存,足够用来存放一个合理大小的激活值,支持正常的 batch size 和序列长度进行训练。
ZeRO Stage 3(最宽松)¶
ZeRO-3最彻底,将模型权重、梯度和优化器状态全部分片。每张卡的基础显存消耗会降至十几GB,留下极大的空间给激活值,可以支持更大的模型、更大的batch size或更长的序列。
结论: $ 8 \times A100\ 80GB $ 训练 13B 模型的显存瓶颈在于优化器状态和梯度,它们是显存消耗的绝对大头。但只要启用了 ZeRO-2 或更高级别的分片策略,就能轻易绕过这个瓶颈,使训练变得可行。

4. 如果模型使用 ZeRO-1,上述估算会如何变化?¶
我们以 13B 模型在 8 张 GPU 上为例,来看看使用 ZeRO-1 分片优化器状态后,每张 GPU 的显存占用会如何变化。
| 显存消耗部分 | 无分片 (DDP) 单卡占用 | ZeRO-1 单卡占用 | 变化情况 |
| 模型权重 (FP16) | 26 GB (完整) | 26 GB (完整) | 不变 |
| 梯度 (FP16) | 26 GB (完整) | 26 GB (完整) | 不变 |
| 优化器状态 (FP32) | 156 GB (完整) | 19.5 GB (分片) | 大幅降低 |
| 激活值 (估算) | ~30 GB | ~30 GB (不变) | 不变 |
| 单卡总显存 | 238 GB (OOM) | ~101.5 GB (仍然OOM) | - |
详细分析:¶
在8张GPU上使用ZeRO-1时,原本每个GPU需要承担的156 GB优化器状态,被平均分配到了8张卡上,因此每张卡只需要存储156/8=19.5 GB。这确实极大地节省了显存。
然而,模型权重(26 GB)和梯度(26 GB)仍然需要每张卡完整保存。因此,即使使用了 ZeRO-1,单卡的基础显存占用依然高达 26 (权重) + 26 (梯度) + 19.5 (分片后的优化器状态) = 71.5 GB。再加上估算的 30 GB 激活值,总需求在 101.5 GB 左右,依然远超单张 A100 80GB 的限制。
这个估算清晰地表明,对于 13B 模型,ZeRO-1 节省的显存并不足以使其在 80GB GPU 上运行。因为此时的主要矛盾已经从优化器状态转移到了模型权重和梯度上。这也是为什么在实践中,必须使用能同时对梯度和优化器状态进行分片的 ZeRO-2,或者连权重也分片的 ZeRO-3,才能解决显存瓶颈。

5. 使用 LoRA 训练 7B 模型(r=8),显存大概是多少?与全量微调对比。¶
LoRA 的核心思想是冻结预训练模型的所有原始权重,并在特定层(如 Attention 的 Q、V 矩阵)旁路注入一组可训练的低秩矩阵。我们以 LLaMA-7B 模型为例,假设使用典型的配置:秩 r=8,仅对 Query 和 Value 投影层添加 LoRA。
可训练参数量估算:¶
LLaMA-7B 有 32 个 Transformer 层,每个层的 Q 和 V 投影矩阵的隐藏维度为 4096。
对于每一个投影矩阵,LoRA 引入一个形状为 $ [4096, 8] $ 的矩阵 A 和一个形状为 $ [8, 4096] $ 的矩阵 B。
每层 Q 和 V 各有两个矩阵(A 和 B),因此每层 LoRA 参数为 $ 2 \times (4096 \times 8 \times 2) \approx 131{,}072 $。总 LoRA 参数量约为 $ 32 \times 131{,}072 \approx 4.2M $(420 万)。这仅仅是原始 7B 参数的极其微小的一部分。
显存对比分析:
| 显存消耗部分 | 全量微调 (7B) | LoRA 微调 (r=8) | 备注 |
| 基座模型权重 | 14 GB (可训练) | 14 GB (冻结) | 冻结后无需梯度和优化器状态 |
| 基座权重梯度 | 14 GB | 0 GB | 冻结,不产生梯度 |
| 基座优化器状态 | 84 GB | 0 GB | 冻结,不产生优化器状态 |
| LoRA 可训练权重 | N/A | ~8.4 MB (FP16) | 参数量极小,可忽略不计 |
| LoRA 权重梯度 | N/A | ~8.4 MB (FP16) | 可忽略不计 |
| LoRA 优化器状态 | N/A | ~33.6 MB (FP32) | 可忽略不计 |
| 激活值 (估算) | ~40 GB | ~40 GB | 两者激活值大小相似 |
| 总显存 | ~152 GB | ~54 GB | 节省约 98 GB |
结论:LoRA之所以能大幅降低显存,是因为它完全消除了与基座模型权重相关的梯度和优化器状态,这两者在全量微调中是显存消耗的绝对主体。将152 GB的总需求降至约54 GB,这意味着以前需要多张顶级GPU才能完成的任务,现在单张A100 (80GB)甚至消费级的48GB显卡都能轻松胜任。
6. QLoRA 的 4-bit 量化底座 + LoRA,显存能降到多少?¶
QLoRA(Quantized LoRA)在 LoRA 的基础上,进一步将冻结的预训练模型基座权重量化为极低精度(例如 4-bit NormalFloat, NF4),从而“压榨”出更多的显存空间。
4 -bit 量化权重的显存占用:¶
对于 7B 模型,FP16 权重需要 14 GB。如果使用 4-bit 量化,理论上每个参数只需 0.5 字节,理论权重大小为 3.5 GB。在实际工程中,加上分组量化所需的额外缩放因子(scale)和零点(zero-point)等开销后,7B 模型的 4-bit 权重通常占用约 3.5 GB 到 4.5 GB 的显存。我们取一个中间值 4 GB。
显存对比分析:¶
| 显存消耗部分 | 标准 LoRA (7B) | QLoRA (4-bit LoRA) | 变化 |
| 基座模型权重 | 14 GB (FP16) | ~4 GB (NF4) | 显著降低 |
| LoRA 可训练参数 | ~8.4 MB (FP16) | ~8.4 MB (FP16) | 不变 |
| 激活值 (估算) | ~40 GB | ~40 GB (不变) | 不变 |
| 总显存 | ~54 GB | ~44 GB | 节省约 10 GB |
结论:QLoRA 将显存占用从标准 LoRA 的约 54 GB 进一步降低到了约 44 GB。虽然和 LoRA 节省的巨量显存相比,这次节省的绝对值不算大,但意义非凡:它将 7B 模型训练的门槛从 48GB 显卡降低到了

7. 训练一个模型时,如何快速估算激活值显存?有什么经验公式?¶
激活值显存没有一个像模型参数那样精确的解析公式,因为它高度依赖于模型架构、批量大小(B)、序列长度(S)和隐藏层维度(H)等运行时参数。不过,我们可以通过分析Transformer内部的中间张量来进行快速估算。
经验估算方法:
在未启用 FlashAttention 和 Gradient Checkpointing 的标准情况下,每层 Transformer 最主要的激活值来自注意力机制。其中,最凸显的是经过 Softmax 之前的注意力得分矩阵,其形状为
[Batch, num_heads, Sequence_Length, Sequence_Length]。
因此,激活值显存与序列长度 SS 的平方成正比,与批大小 BB 成正比。
一个常用的快速估算公式是:
$$ 激活值显存 \ (bytes)\ \approx\ Layers\times\ Batch\times Sequence_{L}ength\times Hidden_{S}ize\times C $$
其中:
Layers:模型的层数。
Hidden_Size:模型的隐藏维度。
C:一个经验常数,通常取15~20字节,涵盖了注意力模块、MLP模块、LayerNorm等所有中间张量。
以 LLaMA-7B 为例:
Layers = 32
Hidden Size = 4096
Batch = 1
Sequence_Length = 2048
取 C = 15 字节
代入公式计算: $ 32 \times 1 \times 2048 \times 4096 \times 15 \text{ bytes} \approx 4.03 \text{ GB} $。
与实际训练中观测到的3~6 GB左右的激活值占用基本吻合。
现代训练中的变数:
启用 FlashAttention:注意力得分矩阵不再被显式存储,激活值显存显著下降,经验常数 C 会变得更小(例如 5~10 字节)。
启用 Gradient Checkpointing:通过牺牲计算时间来换取显存空间,部分中间激活值不再保存,而是在反向传播时重新计算。这可以将激活值显存降低50%甚至更多。
这些优化已成为大模型训练的标配,因此在估算时需要根据具体的配置进行调整。
8. 为什么激活值显存与 batch size 成正比?与序列长度平方成正比?¶
激活值显存的这种特性源于 Transformer 架构中张量的形状和反向传播的需求。
• 与 batch size 成正比¶
在神经网络的前向传播中,每一层的输入、输出以及中间计算结果的第一维通常是 batch_size。例如,输入 token 的嵌入张量形状为 [B, S, H],其中 B 为 batch size,S 为序列长度,H 为隐藏维度。注意力模块的 Query、Key、Value 张量也都是 [B, num_heads, S, head_dim]。
当 B 翻倍时,所有这些张量在每层占用的显存都会翻倍。由于各层独立,总激活值显存自然与 batch size 成正比。这就是为什么调整 batch size 是控制训练显存最直接的手段。
与序列长度平方成正比(无 FlashAttention 时)¶
Transformer 自注意力的核心是计算 Query 和 Key 的点积,得到注意力得分矩阵 $ S = QK^{T} $,其形状为 $ [B, num_heads, S, S] $。这个矩阵必须被保存下来,因为反向传播时需要用它来计算 Q 和 K 的梯度。
该矩阵的元素数量与 $ S^{2} $ 成正比。对于长序列,这个平方项的显存占用成为激活值的大头,远超过其他与 S 成正比的张量(如 Q、K、V)。因此,在标准实现中,激活值显存与序列长度的平方成正比。这就是为什么处理 8K 以上长度的文本时,显存会急剧膨胀。
使用 FlashAttention 后的变化¶
FlashAttention 通过分块计算和在线 softmax 技术,完全避免了对完整 $ S \times S $ 注意力得分矩阵的显式存储。反向传播时通过重计算获得梯度,因此注意力模块的显存复杂度从 $ O(S^2) $ 降为 $ O(S) $。此时激活值显存不再与序列长度平方成正比,而是线性关系。但其他线性层(如 FFN)的激活值仍与 S 成正比,所以整体激活值显存大致与序列长度成正比。
9. 如果序列长度从 2K 扩展到 8K,激活值显存大约增加多少?¶
这取决于是否使用了 FlashAttention。
无 FlashAttention(标准实现)¶
注意力矩阵占据主导,复杂度为 $ O(S^2) $。序列长度从 2048 增加到 8192,变为原来的 4 倍,则注意力矩阵的大小变为原来的 16 倍。虽然其他线性层(如 FFN)的激活值只增加 4 倍,但注意力矩阵通常是激活值中的绝对大头(可占 80% 以上),因此整体激活值显存大约增加 10~16 倍。
例如,2K 时激活值需 4 GB,扩展至 8K 时可能飙升至 40~64 GB,极易导致 OOM。
使用 FlashAttention¶
注意力矩阵不再被存储,激活值主要由线性层产生,复杂度为 $ O(S) $。序列长度扩大4倍,激活值显存大约也扩大4倍(可能稍高,因某些临时缓存仍与 $ S^{2} $ 有关,但 FlashAttention 已将其压缩到最小)。例如,2K 时激活值 4 GB,8K 时约 16~20 GB。
因此,FlashAttention 是长序列训练不可或缺的技术。
10. 假设不使用 FlashAttention,注意力矩阵的显存复杂度是多少?¶
在标准 Transformer 实现中,注意力机制会计算一个得分矩阵(未经过 Softmax):
$$ Scores=Q\times K^{T} $$
其中 Q 和 K 的形状为 [B, num_heads, S, head_dim],点积结果的形状为 [B, num_heads, S, S]。该矩阵在反向传播时必须存在显存中。以 FP16 精度(每个元素 2 字节)为例,其显存占用公式为:
$$ \mathrm{M e m o r y}=B\times\mathrm{n u m_{h} e a d s}\times S\times S\times2\mathrm{b y t e s} $$
时间复杂度: $ O(B \cdot \text{num_heads} \cdot S^{2}) $。这是训练显存中最大的单一项,也是长序列的主要瓶颈。
实例:B=1, num_heads=40, S=8192, FP16.
计算: $ 1 \times 40 \times 8192 \times 8192 \times 2 = 40 \times 67,108,864 \times 2 \approx 5.37 $ GB。
对于有 32 层的模型,如果每一层都存储这样一个矩阵,总激活值将超过 170 GB,这显然不可行。因此FlashAttention 或梯度检查点是必须的。
11. 使用 FlashAttention 后,注意力矩阵显存能节省多少?¶
FlashAttention 通过两种关键技术节省显存:
-
分块计算:将 Q、K、V 分割成小块,逐块加载到 GPU 高速 SRAM 中计算注意力,避免了完整矩阵写入 HBM(显存)。
-
在线 softmax:在分块过程中动态维护归一化因子,无需等待全部得分计算完再做 Softmax。
反向传播时,FlashAttention 不保存注意力矩阵,而是重新计算所需的前向块。因此,注意力矩阵的显存占用从 $ O(B \cdot H \cdot S^2) $ 降为 $ O(B \cdot H \cdot S) $。对于前述 40 头、8192 序列的例子,原本需要约 5.37 GB 的注意力矩阵,使用 FlashAttention 后,这部分显存几乎可忽略(仅保留极小的临时缓冲,通常几十 MB)。对整个模型而言,激活值显存可节省 50%~80%,尤其在长序列时收益巨大。这使得训练 32K 甚至更长上下文成为可能。
12. 梯度检查点一般能减少多少显存?以 Transformer 为例。¶
梯度检查点是一种用时间换空间的策略:前向传播时,仅保存少量“检查点”的激活值,其余中间张量在反向传播时通过重新计算获得。这样,显存中就不必保存全部层的全部激活值。
以标准 Transformer 为例,每层通常需要保存 Q、K、V、注意力得分矩阵、FFN 中间输出等大量张量。如果对整个网络开启梯度检查点(例如每层都设检查点),那么显存中只需保留少数关键张量(例如每层的输入),而其他所有中间结果都可在反向时快速重算。
节省比例:通常可将激活值显存减少50%~80%。具体取决于检查点的放置粒度:
若每1层一个检查点,激活值降至约单层激活值加上一些必要缓存。
若每几层一个检查点,节省略少,但重计算代价也小。
代价:训练速度下降约15%~25%,因为需要额外的前向重计算。在显存紧张时,这是极其有效的妥协。
13. 在训练配置中,如何计算总 batch size 对单卡激活值的影响?¶
在分布式训练中,总 batch size(global batch size)通常被拆分到多个 GPU 上(数据并行),或者通过梯度累积来模拟。单卡激活值仅取决于该卡当前处理的 micro batch size。
- 若增加 GPU 数量:保持 global batch size 不变,单卡 micro batch size 减小,因此单卡激活值显存降低。
若增加 global batch size 而不增加 GPU:单卡 micro batch size 增大,单卡激活值显存增加。
使用梯度累积:假设单卡一次只能处理 micro batch size = 8,但希望模拟 global batch size = 128。此时可以累积 16 步梯度再更新。单卡激活值仍只对应 micro batch size = 8,所以 global batch size 的增大不会直接增加单卡激活值,只是增加了训练时间。
因此,在设置训练配置时,单卡显存瓶颈直接限制了 micro batch size 的上限。
14. 如果模型使用了 SwiGLU 激活,FFN 的参数量和激活值有何变化?对显存的影响?¶
SwiGLU 是 LLaMA 等现代大模型常用的 FFN 结构,相比标准 FFN(两个线性层 + ReLU),它引入了门控机制,增加了参数量和激活值。
· 参数量变化¶
标准 FFN: $ y = W_2 \cdot \text{ReLU}(W_1 x) $。两个权重矩阵,参数个数 $ 2 \times d \times d_{ff} $。
SwiGLU: $ y = (W_3x \odot \text{SiLU}(W_1x)) \cdot W_2 $。有三个权重矩阵,参数个数 $ 3 \times d \times d_{ff} $。通常 $ d_{ff} $ 会相应调整,使总参数量与标准模型持平(例如 LLaMA 中 SwiGLU 的 $ d_{ff} $ 约为标准 FFN 的 2/3),但若保持相同 $ d_{ff} $,则参数量增加 50%。实际工程中,SwiGLU 往往导致总参数量略微增加或持平,但中间激活值显著增加。
激活值变化
前向传播时,除了需要保存 $ W_1 x $ 和 SiLU 输出外,还需保存 $ W_3 x $ 以及两者的逐元素乘积。相比标准 FFN 多出一个线性层输出和一个门控激活值。因此,SwiGLU 的 FFN 激活值比标准 FFN 多出约 50%。
对显存的影响¶
权重显存:若参数量增加,权重显存同比例增加。
激活值显存:FFN部分的激活值增加约50%,由于FFN激活值在总激活值中占比可观,整体激活值显存可能增加20%~30%。
因此,SwiGLU 在提升性能的同时,也增加了显存压力,需要配合梯度检查点等技术使用。
15. 训练一个 MoE 模型(如 Mixtral 8x7B),总参数量 47B,训练时显存应如何估算?¶
MoE(Mixture of Experts)模型的特点是总参数量巨大,但每个 token 只激活部分专家(Mixtral 8x7B每次激活2个专家)。训练显存的估算必须考虑所有专家的权重、梯度和优化器状态,因为这些都需要在训练时驻留显存(或通过分片策略分布到多卡)。
估算示例(Mixtral 8x7B,FP16 + AdamW):
• 权重显存(FP16):总参数量 47B,需 $ 47 \times 10^9 \times 2 $ bytes = 94 GB。
梯度显存(FP16):训练中需为所有参数分配梯度空间,同样94 GB(即使只有激活的专家有非零梯度,但框架通常为全部参数分配)。
- 优化器状态(FP32):AdamW 需要主权重副本(4B/param)、动量 m(4B)、动量 v(4B),共 12 字节/参数。总优化器状态 $ 47 \times 10^9 \times 12 $ bytes = 564 GB。这是显存的最大消耗者。
激活值显存:由于每次只激活2个专家(约13B等效稠密模型),激活值与一个13B稠密模型相当,通常在40~60 GB(取决于batch size和序列长度)。
总显存需求(单卡理论上限): $ 94+94+564+50\approx802\ GB $。显然无法单卡运行,必须采用专家并行+ZeRO分片。
分布式策略:¶
专家并行:将8个专家分到8张GPU上,每张卡只存储自己负责的专家权重、梯度和优化器状态。单卡权重94/8≈12GB,梯度12GB,优化器状态564/8≈70.5GB。合计基础开销约94.5GB,加上激活值,单卡依然紧张。因此通常还需要结合ZeRO-2或ZeRO-3,进一步分片梯度或优化器状态,或者使用CPU offload将部分优化器状态卸载到内存。
16. MoE 模型所有专家的参数都必须常驻显存吗?¶
训练时:是的,原则上所有专家的参数(以及它们对应的优化器状态)都需要能够被访问,因为优化器需要为每个参数维护动量,且在反向传播时被激活的专家需要更新。不过,可以通过 ZeRO-Infinity 等 offload 技术将部分不活跃的专家参数或优化器状态临时卸载到 CPU 内存甚至 NVMe,但速度会严重下降。在标准的 GPU 训练中,为追求效率,通常将所有专家参数放在显存中,然后通过专家并行分散到多卡。
推理时:不一定。推理只需前向,不需要梯度和优化器状态。如果单卡显存紧张,可以采用专家卸载(Expert Offloading):将不常用的专家权重放在CPU内存,当某个专家被激活时再加载到GPU。或者使用专家并行,每张GPU只驻留部分专家,通过All-to-All通信传递token。这样可以大幅降低单卡显存需求,代价是延迟增加。
17. 专家并行(Expert Parallelism)如何减少单卡显存?¶
专家并行是专门针对 MoE 模型设计的分布式策略。它将不同的专家放置在不同的 GPU 上。例如,Mixtral 8x7B 有 8 个专家,使用 8 张 GPU 时,可以每张卡负责 1 个专家。
显存减少机制:¶
权重分片:每张 GPU 只存储自己负责的专家的权重,而非全部 8 个专家。单卡权重显存降为总权重的 1/(专家并行度)。
优化器状态分片:同理,每张 GPU 只需为自己负责的专家维护优化器状态,优化器状态显存也降为 1/(专家并行度)。这极大缓解了 MoE 训练中优化器状态爆炸的问题。
激活值:激活值仅与本次被激活的专家有关,大小相当于一个等效稠密模型。专家并行不会减少单卡的激活值(实际上 token 路由可能导致某些卡上的激活专家数略多),但也不会显著增加。
专家并行通常与数据并行、ZeRO结合,实现高效的 MoE 训练。例如,总 32 张 GPU,可以分成 4 个数据并行组,每组内 8 张 GPU 做专家并行。这样每张卡既享受了数据并行的梯度同步,又通过专家并行分摊了庞大的参数和优化器状态。