跳转至

训练时显存构成

一、训练一个模型时,显存里都存放了哪些主要数据?

当我们启动一个分布式训练任务时,GPU HBM 显存中存放的数据可以精确划分为以下 六大门类。任何显存不足(OOM)问题,都可以归因于它们中的一项或多项失控。

image.png

深层剖析: 实际框架中还存在内存碎片。PyTorch 的缓存分配器会预留大块显存,用过的空间不立即释放,这可能导致 nvidia-smi 看到的占用高于理论计算,需使用 torch.cuda.memory_allocated()torch.cuda.memory_reserved() 区分。


二、模型权重在显存中存储的精度有哪些?各占多少字节?

这不仅是字节的问题,更是动态范围、精度与硬件架构的深度耦合。面试官想听你说出 IEEE 754 标准细节,并联系实际训练。

常见浮点格式细目

查看内嵌表格

深入理解混合精度训练中的“多份权重”

标准 FP16 混合精度训练 的显存中实际存在:

  • FP16 权重 (2Φ):用于前向计算,减少访存和显存占用;

  • FP32 主权重 (4Φ):优化器用 FP32 状态更新参数时,需要 FP32 精度累加避免舍入误差。每个 step 将 FP16 权重同步为主权重的截断版。

若启用 BF16,则 BF16 权重直接用作计算,但仍需 FP32 主权重用于更新,因为 BF16 的 7 位尾数不足以保证微小的权重更新累积。


三、为什么训练时的梯度大小与模型权重完全相等?

这个问题看似简单,但深入回答需要从自动微分原理和计算图入手。

核心原因:参数与梯度的“对偶性”

假设模型有一个参数 W,损失 LW 的函数。梯度定义为:

image.png

链式法则要求:每一个可学习的张量 W 都必然接收从下游传回的梯度信号。这个梯度信号必须具有与 WW 完全相同的形状,因为更新时执行:

image.png

是逐元素的运算。

深入自动微分实现

在 PyTorch 的计算图中,每个张量有一个 .grad 属性。当调用 loss.backward() 时,每个 requires_grad=True 的叶子张量都会累加所有路径传来的梯度。框架分配的存储与原始张量形状严格一致,不可分块或压缩。

image.png

常见误区

  • “如果参数共享呢?” 梯度是按张量单元累加的,如 embedding 共享权重,其梯度会在反向时累加,但最终也只有一个与原始权重同形的张量,显存占用不变。

  • “稀疏梯度可以变小。” 对,框架支持稀疏梯度,但存储格式不同,但展开成稠密张量时依然占用同样元素数,但稀疏可节省显存。


四、Adam 优化器需要保存哪些状态?各占多少显存?

Adam 状态细节

标准 Adam 更新公式:

image.png

为每个参数 θ 保存:

  • m (一阶矩):形状与 θ 相同,通常 FP32。

  • v (二阶矩):形状与 θ 相同,通常 FP32。

  • 某些实现保存步数 t:一个标量,可忽略。

因此,每个参数对应 8 字节(FP32 下 4+4)。

扩展:AdamW 与 8-bit Adam

  • AdamW:权重衰减与梯度更新解耦,但 m,v 不变,状态量完全相同。

  • 8-bit Adam (如 bitsandbytes):将 mv 量化为 8 位,每参数平均 1 字节,再配合量化 block 的常数存储,平均每参数约 2 字节,显存降低至 1/4。

image.png


五、为什么 Adam 优化器状态通常用 FP32 存储?如果用 FP16 会怎样?

这需要从 FP16 的动态范围和 Adam 的计算特性深度分析。

FP16 的下溢灾难

image.png

FP16 的舍入误差累积

image.png

损失缩放 (Loss Scaling) 治标不治本

混合精度训练用放大损失来将小梯度推入 FP16 范围,然后反向缩放权重更新。但这只能缓解梯度本身的溢出,无法改变优化器状态的更新过程——因为 m,v 的更新仍在优化器内部,使用的是原始梯度(缩放后的梯度除以缩放因子,又变回小值)。所以必须用 FP32 存储状态。

BF16 可行吗?

BF16 动态范围与 FP32 相同(8位指数),不会下溢。但其 7 位尾数导致累加精度不足。实验表明直接用 BF16 存 m,v 会使收敛不稳定,因此依然推荐 FP32 状态。

面试金句:“Adam 状态是随机优化过程的长期记忆,需要高精度存储以防止缓慢累积中的信息丢失;而前向激活可以接受低精度,因为它在每一步都是重新计算的。”


六、激活值是什么?为什么它会在显存中占据大量空间?

定义:不仅仅是“中间结果”

激活值(Activation)是指前向传播中所有需要在反向传播中重用的张量,以满足链式法则的需要。包括但不限于:

  • 每层的输入和输出(需用于梯度传播)

  • 非线性函数的输入(如 ReLU 需要知道哪些元素>0)

  • 归一化层的输入、均值、方差(用于梯度计算)

  • Dropout 的 mask

  • 注意力机制中的 QKT 矩阵、softmax 后的概率矩阵等

为什么它是个巨大的内存消费者?

  1. 保存全量前向状态 默认的自动微分采用“全检查点”策略:计算图创建时,每个算子的输入被保存。对于一个 L 层的 Transformer,每次 forward 会生成并保留 L 个层的完整状态,就像给每一层都拍了张快照。

  2. 微观尺寸不可小觑

image.png

  1. 无法被压缩 这些张量必须保持高精度,因为反向传播需要准确的数值计算梯度。虽然可以用梯度检查点(Activation Checkpointing)丢弃大部分,只保留少量重计算所需,但默认配置下,激活值往往占据总显存的 30%~50%。

激活显存总量的估算公式

image.png


七、对于一个 Transformer 层,哪些激活值占用显存最大?

我们需要深入一个标准 Transformer 块的内部,按前向顺序列举并计算各激活的大小。

image.png

层内激活组成

image.png

FlashAttention 的拯救

FlashAttention 将 QKT 分块计算,不将完整的 T×T 矩阵驻留在显存中,仅保留 softmax 的归一化统计量。这使最大激活变为 FFN 中间层和注意力输出,激活峰值降低一个数量级。


八、为什么激活值显存与序列长度有关?关系是怎样的?

序列长度 TT 决定激活有两种依赖性,必须拆开理解。

线性依赖 O(T)

所有与 token 数量成线性关系的张量:隐状态、Q、K、V、FFN 输入输出等。它们形状包含 T 维度,即每个 token 贡献一组特征。这类占用随 T 线性增长,斜率由 BD 决定。

image.png

image.png

九、训练的显存占用中,通常权重、梯度和优化器状态哪个占比最大?何时改变?

这需要放在典型配置下看,以及不同优化器和并行策略下的动态对比。

标准场景:混合精度 Adam + 全参数训练

数据关系(以每参数字节计):

  • FP16 计算权重:2

  • FP32 主权重副本:4

  • FP16 梯度:2 (假设保持 FP16)

  • FP32 Adam 状态 (m,v):8

  • 总基础静态内存 = 16 字节/参数

饼图分析:

  • 优化器状态:8/16 = 50%

  • 主权重:4/16 = 25%

  • 计算权重+梯度共:4/16 = 25%

优化器状态始终最大,且稳定不变。

何时优化器不再是第一大项?

查看内嵌表格

结论:对于现代 LLM 的全参数 Adam 训练,优化器状态是毋庸置疑的第一显存消耗者,这也是 ZeRO-1 将优化器状态分片能立竿见影节省 4 倍显存的原因。


十、临时缓冲区和通信缓冲区是什么?它们占用显存大吗?

通常不会成为瓶颈,但在特定极端配置下可能引发 OOM。

临时缓冲区 (Workspace)

深度学习的底层库(cuBLAS, cuDNN, 自定义 CUDA kernel)在执行算子时需要 scratch memory:

  • 矩阵乘法分段:大矩阵乘法可能会划分 tile,临时存储部分积。

  • Softmax / LayerNorm:需要存储最大值、求和等中间统计量。

  • 注意力融合:如 FlashAttention 需要在 SRAM 和 HBM 之间搬运数据,需要额外预留缓冲区用于分块计算的重组。

  • Dropout:需要存储随机 mask,但通常与激活耦合。

大小:一般单个算子几 MB 到几十 MB,PyTorch 的 cublas 句柄维护一个共享 workspace。在 ResNet 等 CNN 里 cuDNN 可能预先分配数百 MB。对于 Transformer,单次迭代的临时缓冲一般 < 2GB,相对权重和激活是零头。

通信缓冲区

分布式通信库(NCCL)执行 AllReduce 或 AllGather 时,需要将数据打包成连续块。

  • Flatten Buffer:早先框架会将所有梯度摊平进一个大 buffer 进行 AllReduce,此 buffer 大小等于梯度总大小(≈2Φ)。对 100B 模型,就是 200GB,这太恐怖。

  • Bucket 策略 (PyTorch DDP):将参数/梯度分桶,每次通信一个桶(典型 25~128MB)。只需要每个桶大小的缓冲。显存峰值可控。

  • ZeRO 的通信:在 ZeRO-3 中,做 allgather 收集参数分片时,需要为每层分配临时 buffer 容纳全量参数。通常通过流水线预取隐藏,峰值一般为单层参数大小 + 某个 bucket。

  • 序列并行 / 张量并行:层间需要插入 allreduce,需为每个并行通信分配临时 buffer,通常不大。

潜在风险:当使用超大 batch 且梯度累积步数较多时,梯度的 allreduce 可能设计为一次性通信(未分桶)导致超大 buffer 分配,或者某些自定义融合算子临时缓冲区设计不当。务必用 torch.cuda.memory_summary()nsys 分析。


十一、框架(如 PyTorch)自身的显存开销大约占多少?由什么构成?

“框架自身开销”通常指除模型、数据、优化器之外,由深度学习框架运行时本身占用的、无法被用户直接控制的 GPU 显存。它不是某个可以单独列出的“项目”,而是由多个幕后机制叠加形成的隐性消耗。对于现代 LLM 训练来说,这部分开销通常在 数百 MB 到几 GB 之间,相对于数百 GB 的参数和激活来说占比不大(<5%),但足以成为压垮 OOM 的最后一根稻草。

其构成主要包括以下几个方面:

CUDA 上下文与驱动开销

  • 每个 CUDA 设备初始化时会分配一个主上下文,约 100–300 MB。

  • 包含设备内存管理表、内核模块、统一内存句柄等。

  • 这是进程级别的固定开销,与模型大小无关。

cuBLAS / cuDNN 句柄及其 Workspace

  • PyTorch 会为每个设备维护一个 cuBLAS 句柄和 cuDNN 句柄,并附设一个共享 workspace 缓冲区。

  • 在卷积网络中,cuDNN 为寻找最优算法可能需要分配数十 MB 甚至几百 MB 的临时空间;Transformer 中矩阵乘法亦会占用。

  • PyTorch 默认会预先分配一个 workspace(可通过环境变量 CUDNN_WORKSPACE_LIMIT 等限制),以避免运行时频繁申请。

  • 典型大小:几十 MB 到 几百 MB。

内存分配器的缓存碎片

  • PyTorch 使用缓存内存分配器(Caching Allocator)。它不会在每次张量释放后立即将显存还给 CUDA,而是保留为 cached memory,以待下次分配。

  • 这导致 nvidia-smi 显示的显存占用远高于实际张量占用。

  • 由于分配器会产生碎片(内存块分割),可能导致即使空闲总字节足够,却无法分配出一个连续的大张量,从而触发 OOM。

  • 碎片本身不直接算作“开销”,但不可用的碎片空间在效果上等同于被框架占用。

Autograd 计算图元数据

  • 每个需要梯度的张量都会在 CPU 和 GPU 上维护一个 Node 对象(即计算图节点),记录算子类型、反向函数、边关系等。

  • 图节点本身占用 CPU 内存为主,但部分反向传播需要的中间状态(如保存的张量列表)会保留在显存中。

  • 对于百万级的小模型,计算图开销极小;对于有数万个小张量参与的大型图,累积的元数据可能达到 数十 MB。

通信库的持久化缓冲

  • NCCL(NVIDIA Collective Communications Library)初始化时会建立通信环,并为每个 ring 分配一小块持久化缓冲用于控制信息。

  • 此外,NCCL 可能会为一些内部操作申请少量显存。

  • 一般 数十 MB 量级。

PyTorch 内部的全局状态和分配器

  • 各种 CUDA 流(stream)、事件(event)、以及 C++ 侧对象的内存。

  • 使用 torch.cuda.memory_summary() 可以观察到一些静态分配。

量化总结:

  • 在干净启动 PyTorch 并加载大模型后,除了张量数据外,框架自身开销约为 0.5–2 GB。

  • 影响因素:GPU 代数、CUDA 版本、cuDNN 版本、是否开启 benchmark 模式、NCCL 配置等。

  • 面试要点:当发现显存离预期只差几百 MB 时,可以尝试通过清理缓存、减少 workspace 限制、使用 torch.cuda.empty_cache() 等方法挤出空间,这其实就是在挤压框架的缓存部分。


十二、PyTorch 的 memory_allocated()memory_reserved() 分别代表什么?

理解这两个函数是 PyTorch 显存调试的核心,它们精确刻画了“你用了多少”与“框架占了但没还”的区别。

1. torch.cuda.memory_allocated(device)

  • 含义:当前在该设备上所有活跃的张量(以及某些内部存储)实际占用的显存字节数。这是用户可见的、确切被计算图、模型、数据使用的空间。

  • 变动方式:

  • 创建张量时增加;
  • 张量失去引用并被 del 或作用域结束后,经过 Python GC 调用,对应的 CUDA 存储被标记为释放,allocated 减少。

  • 特点:严格反映“正在使用”的显存,等同于 nvidia-smi 中的“used”减去缓存的空闲块。

2. torch.cuda.memory_reserved(device)

  • 含义:PyTorch 缓存分配器从 CUDA 驱动申请并保留的总显存字节数。它包括 allocated 空间,再加上已释放但尚未归还给操作系统的显存块(即缓存)。

  • 关系:reserved >= allocated。差值代表“空闲但被分配器持有的显存”,可用于快速满足未来的张量分配,而无需昂贵的 CUDA API 调用。

  • 变动方式:

  • allocated 增加时 reserved 可能也需向 CUDA 申请新块而增加;
  • 张量释放时,allocated 减少,但 reserved 不减,除非显式调用 torch.cuda.empty_cache()

  • 目的:减少 cudaMalloc/cudaFree 的调用次数,提升性能。

实战解读

运行 torch.cuda.memory_summary() 会输出详细的统计表,包括:

  • Allocated memory: 当前活跃字节

  • Reserved memory: 缓存持有字节

  • Active memory: 活跃字节(同 allocated)

  • Inactive memory: 被释放但仍在缓存中的字节(即 reserved - allocated

面试时的精辟总结:

  • allocated 是你“正在吃”的饭;

  • reserved 是“你碗里 + 锅里已经盛好但还没吃”的饭;

  • 如果你看到 nvidia-smi 占用远高于 allocated,很可能是大量 inactive 缓存未释放,或框架其他部分(如 cuDNN workspace)直接向 CUDA 申请了显存而绕过了 PyTorch 分配器。


十三、为什么训练时显存占用会波动?哪个阶段是峰值?

显存占用在单个训练迭代内不是恒定的,而是随着前向、反向、优化器更新的推进,呈现周期性波峰和波谷。这种波动性是显存优化的关键突破口。

波动的主要来源

  1. 激活值的生灭
  2. 前向传播:逐层计算并保留激活张量,显存持续攀升,直到最后一层完成前向。
  3. 反向传播:自动释放不再需要的激活(例如已经用过的中间结果),同时产生梯度张量。
  4. 因此,前向结束到反向开始之间,激活量达到峰值;反向过程中激活逐步释放,而梯度逐渐累加,整体显存可能先升后降。

  5. 梯度累积与 AllReduce

  6. 在 DDP 中,反向传播会异步触发梯度 AllReduce。通信期间,梯度缓冲区被占用。
  7. 如果使用了梯度累积(gradient accumulation),每一步 backward() 会累加梯度,梯度显存持续存在,峰值更高。

  8. 临时缓冲区(Workspace)

  9. 某些算子(如矩阵乘法、注意力融合)会临时申请大块 workspace,计算完成后立即释放。这会在极短时间内推高显存水位。
  10. 例如,某个 kernel 临时需要 1 GB 的 scratchpad,使用完毕立即释放。

  11. 优化器步骤的临时分配

  12. optimizer.step() 中,优化器可能会创建临时张量(例如 FP32 梯度的拷贝、缩放计算),然后更新状态。虽然通常不大,但叠加在已经很高的水位上可能成为最后的一击。

显存峰值常出现在两个阶段:

场景一:标准 DDP 全参数训练(无检查点)

  • 峰值阶段:前向刚完成、反向尚未开始或刚开始时。

  • 原因:所有层的激活值全部留存,等待反向使用。此时激活值总量最大,加上权重、优化器状态等静态数据,达到整个迭代的最高水位。

场景二:启用梯度检查点(Activation Checkpointing)

  • 峰值阶段:反向传播到某个检查段落的末尾,此时需要重新前向计算这一段落的激活,会短暂地重新产生该段的激活值,与当前已保留的其他激活值叠加。

  • 峰值可能不是全局最大,而是局部“尖峰”,但依然高于稳态。

场景三:通信与计算重叠

  • 在 Nvidia 的 allreduce 分桶策略中,梯度桶的通信可能会和反向计算重叠,导致梯度的通信缓冲区和激活值同时存在,推高瞬时峰值。

实战经验:

  • 观察 nvidia-smi 或利用 PyTorch profiler 的显存时间线,会看到类似“鲨鱼鳍”形状:前向攀升,反向缓慢下降。

  • 避免峰值 OOM 的技巧:启用梯度检查点(牺牲计算换显存)、减小 micro batch size、使用 ZeRO 分片,或者在反向传播开始前手动删除不再需要的参考。


十四、在使用梯度检查点时,激活值显存会发生什么变化?

梯度检查点(Gradient Checkpointing,也称作 Activation Checkpointing)是一种用计算换显存的技术,它以增加一次额外的子图重前向计算为代价,将激活值显存压缩到亚线性甚至常数量级。

默认行为的激活存储模型

image.png

检查点的核心思想

  • 将计算图划分为多个检查段(segments)。PyTorch 的 torch.utils.checkpoint 默认以一个 nn.Module 为边界(也可以自定义)。

  • 前向阶段:只保留检查段输入张量,段内产生的中间激活全部丢弃(不保存)。

  • 反向阶段:当需要该段的梯度时,利用保存的输入,重新执行一次段的前向计算,临时恢复所需的中间激活,计算完该段的所有梯度后,再将这些临时激活丢弃。

  • 这样,任何一个时刻,只需要存储少量的检查点输入和当前段的临时激活,而无需保存所有层的全部激活。

激活显存的变化(按段划分)

假设一个模型有 L 层,划分为 k 个段(每段包含 L/k 层)。

  • 无检查点:激活量 ∝L

  • 全部检查点(每层一个段):激活量 ∝1(只需一个层的激活临时重建,加上检查点输入的少量开销)。但增加了 33% 的计算(每个层都重算一次前向)。

image.png

需要注意的副作用

  1. 计算开销:每段多执行一次前向,总的额外计算量约为 33%(如果每层都重算)。可通过选择性检查点(仅对注意力部分检查点,保留 FFN 激活)来降低开销。

  2. 与反向传播的结合:重算的前向仍然需要临时显存来存放该段的中间结果,峰值可能比无检查点的平稳水位更高,形成局部尖峰。

  3. 随机操作:重算过程中 Dropout 等的随机性必须被固定(通过保存 RNG 状态),否则重算结果与原始前向不一致。

一句话总结:梯度检查点把激活值的常驻显存从整座山压缩成一张薄片,只在反向时短暂地“租用”一片空间进行计算,随后立即归还。


十五、模型并行(TP/PP)是如何影响单卡上的显存构成的?

当模型大到单卡放不下,我们需要用张量并行(TP)和流水线并行(PP)将模型切分到多卡。二者对显存的影响机制完全不同。

张量并行(Tensor Parallelism, TP)

TP 将单个层内的参数矩阵在多个 GPU 上切分,每张卡只持有参数的一部分。

对显存各类别的影响:

  • 权重:每张卡只持分片后的权重,例如将 D×D 的矩阵按列切分成 D×(D/N),单卡权重大幅减小(约 1/N)。

  • 梯度:同样,梯度与权重分片对应,单卡梯度也降至 1/N

  • 优化器状态:由于每张卡只更新自己的那部分参数,mm 和 vv 也相应减少到 1/N1/N

  • 激活值:情况复杂。TP 的前向计算中,各卡并行计算自己的分片,但某些操作需要在同层内进行AllReduce 通信。激活值(如隐状态)通常会在 TP 组内复制或分片。具体:

  • Megatron 风格 TP:在前向时,每张卡保留自己的分片输出,但注意力需要 allreduceallgather。激活张量的形状在 TP 维度上减少,但可能需要额外的通信缓冲区。
  • 有效激活存储通常也与 1/N 有关,但因为通信重叠,可能并不严格线性。

  • 通信缓冲区:新增了同步所需的临时缓冲,大小与分片大小相关。

总体:TP 将几乎所有静态显存(权重、梯度、状态)等比压缩,激活部分通常也减小,是一种“全面瘦身”。但代价是通信量和频率极高,需高速互联(如 NVLink)。

流水线并行(Pipeline Parallelism, PP)

PP 将模型按层切分,每个 GPU 负责一段连续的层(例如 GPU 0 负责层 1-10,GPU 1 负责层 11-20)。

对显存各类别的影响:

  • 权重、梯度、优化器状态:每张卡只持有自己那几层的参数,故这些部分也等比减小(层数减少为 1/PP)。

  • 激活值:最核心的不同。由于 PP 一次只在一个微批次(micro-batch)上经过该卡负责的层,所以激活值只包含这几层的中间结果,而不是全部层。激活值同样减少了约 1/PP 倍(基于层数比例)。

  • 额外的内存:PP 通常需要将多个微批次的中间激活(跨阶段的边界张量)暂存起来,以便反向时使用。这是流水线气泡引入的额外激活,但占用通常和段内激活相当,不会显著增加峰值。

  • 通信缓冲区:PP 的点对点通信(发送/接收激活和梯度)需要临时缓冲区,大小等于一个微批次的激活张量。

总体:PP 同样减小了各静态项和激活项,但带来新的激活暂存需求。它可以与 TP 正交组合(3D 并行),共同降低单卡总显存。

组合效果

3D 并行下,每个 GPU 只持有模型的一个小碎片:

  • 参数 = Φ/(TP×PP)

  • 优化器状态同等缩减。

  • 激活值 = Lper_gpu× 单层激活,与 PP 成反比,且 TP 可能进一步减小单层激活大小。

  • 这实现了千卡训练超大模型的可能性。

面试者应强调:TP 减小了每个张量的大小(切割矩阵),而 PP 减小了张量的数量(切割层数),两者正交叠加使得单卡负担成倍缩小。


十六、混合精度训练中,除了 FP16 的权重和梯度,为什么还有一个 FP32 的主权重副本?

在混合精度训练(例如 torch.cuda.amp)的标准实现中,虽然前向和反向使用半精度(FP16/BF16)来加速,但必须额外维护一份 FP32 的主权重(Master Weights)。这不是可选项,而是数值稳定性与精度要求下的工程必然。

原因一:权重更新的累加精度需求

image.png

原因二:梯度缩放(Loss Scaling)只是权宜之计

image.png

原因三:Adam 等自适应优化器的状态是 FP32,需要 FP32 权重参与

Adam 内部的 m、v 以 FP32 存储,它们累积的是 FP32 梯度(从 FP16 反量化或直接转换)。更新时计算出 FP32 的 Δw,自然应该应用到 FP32 的权重上。

工作流

  1. 训练开始:FP32 主权重初始化,并将其转换为 FP16(截断)用于前向计算。

  2. 前向:使用 FP16 权重和 FP16 激活计算。

  3. 反向:得到 FP16 梯度。

  4. 优化器内部:将 FP16 梯度转换为 FP32(如果混合精度代码如此),用于更新 FP32 的 m,vm,v

  5. 计算 FP32 的 Δw,更新 FP32 主权重。

  6. 将更新后的 FP32 主权重再次转换为 FP16,供下一步前向使用。

因此,显存中同时存在:FP16 计算用权重 + FP32 主权重 + FP16 梯度(如果保留) + FP32 状态。FP32 主权重带来的额外 4 字节/参数,是保证模型收敛的代价。


十七、解释 ZeRO-1 如何优化显存占用?具体省了哪部分?

ZeRO(Zero Redundancy Optimizer)是 DeepSpeed 提出的分布式训练显存优化系列。ZeRO Stage 1 的核心思想是:在数据并行组内,将优化器状态(m 和 v)分片(partition)到各 GPU,而不是每张卡保存完整的全局优化器状态。 同时,每张卡仍然持有完整的模型权重和梯度(用于各自的数据批次),但优化器状态不再冗余。

原始 DDP 的显存状态

在标准数据并行中,每个 GPU 都有:

  • 一份完整的模型权重(FP16 + FP32 主副本)

  • 一份完整的梯度(用于各自 micro-batch 的累加)

  • 一份完整的优化器状态(m 和 v,FP32)——这部分在所有 GPU 之间完全相同,是冗余。

因此,总的跨所有 GPU 的优化器状态存储为 N×8Φ 字节,其中 Φ 是参数量,N 是 GPU 数量。

ZeRO-1 如何消除冗余

  • 在反向传播之后、优化器更新之前,所有 GPU 执行一次 Reduce-Scatter 通信,将各卡持有的完整梯度进行求和,同时各卡只保留与自己分片对应的梯度部分。

  • 这样,每张 GPU 只持有一部分梯度的最终平均结果(即全局梯度的分片)。

  • 然后,每张卡仅对自己持有的那一部分参数,更新其对应的优化器状态 m,v 和 FP32 主权重。

  • 更新完成后,执行 All-Gather,将更新后的参数分片从各个 GPU 收集回完整的参数,供下一步前向使用。

显存节省:

  • 优化器状态从每卡的 8Φ 降至 8Φ/N(精确分片)。

  • 省了哪部分?:省去了 N−1 份冗余的优化器状态。对于 N=8 的情况,优化器状态显存降低为原来的 1/8。

  • 权重和梯度的存储仍然是每卡完整持有(FP16 权重及梯度),因此它们不省。

通信代价

  • 新增一次 Reduce-Scatter(梯度求和与分片)和一次 All-Gather(更新后参数收集),通信量与标准 DDP 的 AllReduce 梯度相当。总体通信量大约为 DDP 的 1.5 倍,但换来显著的显存节省。

一句话:ZeRO-1 只切优化器状态,不切权重和梯度,用少量额外通信打破了优化器状态的每卡冗余。


十八、ZeRO-2 和 ZeRO-3 又分别省了什么?

ZeRO 的三个阶段逐步递进,切分越来越彻底。

ZeRO Stage 2(优化器状态 + 梯度分片)

  • 额外分片:梯度。

  • 在 ZeRO-1 的基础上,每一张卡也不再保存完整的梯度。反向传播过程中,每层的梯度计算出后,立即对这部分梯度执行 Reduce-Scatter,使得每张卡只保留与自己参数分片对应的梯度分片,而非等全部梯度算完再一次性通信。

  • 这样,梯度所需的显存也从每卡 2Φ(FP16)降至 2Φ/N

  • 省了哪部分?:在省掉冗余优化器状态的基础上,又省掉了冗余的梯度。

  • 显存节省效果:每卡只需存储自己负责更新的那部分参数的梯度,大大缓解大型模型在较多 GPU 上的梯度瓶颈。

ZeRO Stage 3(优化器状态 + 梯度 + 参数分片)

  • 进一步分片:模型参数本身。

  • 在 ZeRO-2 的基础上,连模型权重也不在每张卡上保留完整副本。每张卡只持久化保存与自己的参数分片对应的那一部分 FP16 权重和 FP32 主权重副本。

  • 前向传播时,需要哪层参数,就通过 All-Gather 从其他 GPU 收集该层的完整参数,计算完即丢弃(或保留少量缓存)。

  • 反向传播同样:重算参数或收集参数来计算梯度。

  • 省了哪部分?:参数存储从每卡 2Φ(FP16)+ 4Φ(FP32 主副本)降为 (2+4)Φ/N

  • 这是最彻底的显存优化:优化器状态、梯度、参数全部分布式存储,单卡总显存大约降至原始的 1/N,从而实现千亿参数模型在几十张卡上的训练。

通信开销比较

  • ZeRO-3 在每层前向和反向都需要频繁 All-Gather 参数分片,通信量比 ZeRO-2 显著增大(约 1.5–2 倍)。因此需要极高的节点内带宽(NVLink、InfiniBand)来保持计算不被通信阻塞。

  • 但它带来的显存收益允许 batch size 更大或模型更大,整体吞吐量可能更高。

总结对比如下:

查看内嵌表格


十九、使用数据并行(DDP)时,每张卡的优化器状态是冗余的吗?

是的,完全冗余。 这是 DDP 显存效率低下的一个重要原因,也是 ZeRO 诞生的初衷。

DDP 的工作机制

  • 每个 GPU 拥有一个完整的模型副本,以及独立的优化器实例。

  • 前向:各卡使用自己分配的 micro-batch 数据,计算出各自的 loss。

  • 反向:各卡计算各自的梯度。

  • 梯度同步:在所有 GPU 间执行 AllReduce,平均梯度,使所有 GPU 最终持有完全相同的全局平均梯度。

  • 优化器更新:每张卡独立地用相同的梯度更新自己的模型和优化器状态。

  • 结果是所有 GPU 的参数和优化器状态始终保持一致。

冗余的本质

因为所有 GPU 都执行一模一样的参数更新(更新量相同),优化器状态 m 和 v 在更新后会完全相等。这些完全重复的状态在 N 张卡上总计占用 N×8Φ 字节,但实际只需 8Φ 的唯一信息。 因此,对于 NN-way DDP,优化器状态是 N 倍冗余。

为什么还用它?

  • 实现极其简单,通信只有 AllReduce 梯度,没有额外的状态分片管理。

  • 对于中小模型(参数 < 1B),单卡放得下完整优化器状态时,DDP 的简洁性胜出。

  • 显存浪费在模型较小时可以接受(8 GB 冗余对于8卡,每卡才多1 GB,影响不大)。

面试要点:当面试官问“DDP 中优化器状态是不是浪费”,你要明确指出:是信息上的完全冗余,但它是为了保持每卡独立更新并简化通信而付出的代价;ZeRO 正是通过打破这一冗余来支撑大模型训练。


二十、训练中 batch size 增大对显存的影响主要体现在哪一部分?

增大 batch size 对显存的冲击是高度非对称的,主要集中在激活值部分,而对权重、梯度、优化器状态几乎没有影响(在使用数据并行或 ZeRO 的情况下需细究)。

对激活值的影响(最大头)

  • 前向传播产生的大部分中间张量的第一个维度是batch size(即 micro-batch size),例如隐状态 [B,T,D]、注意力矩阵 [B,h,T,T]。

  • 因此,激活值总量与 BB 呈正比(忽略 T2 项时)。如果 BB 翻倍,激活值显存大约也翻倍(注意 T2 项也会因为多头而随 B 倍增)。

  • 这是导致 OOM 的最常见原因:加大 batch 想提高吞吐,结果激活瞬间爆掉。

对梯度和优化器状态的影响

  • 在数据并行(DDP)下,每个 GPU 仍然保留完整的梯度和完整的优化器状态,这些大小与 batch size 无关,只与模型参数量有关。因为梯度张量的形状和参数一样,不随输入 batch 变化。

  • 即使使用梯度累积(多个 micro-batch 累加梯度),梯度存储依然是固定的 P 大小,不会因累积步数而变大。

  • 因此,梯度与优化器状态的显存对 B 是免疫的。

对权重的影响

  • 同样,权重与 batch size 无关。

特殊场景

  • 当使用梯度检查点时,激活量虽被压缩,但仍有一个和 B 成正比的基本激活分量(检查点输入)。增大 B 依然会线性增加这一分量。

  • 在模型并行/ZeRO-3中,参数是分布式存储,但激活量依然是每卡的局部批量导致的,未被打折。因此大 batch 训练需要足够多的显存来容纳激活,或者结合张量并行来降低每卡的激活大小。

如何应对 batch 增大导致的激活膨胀?

image.png