跳转至

DeepSpeed Memory and Parallelism

导言

“显存不够”不是一个足够精确的诊断。可能是优化器状态常驻 GPU,可能是 ZeRO-3 跨节点通信暴露,也可能是单层矩阵本身无法放进一张卡。ZeRO-Offload、ZeRO++、MixZ++ 和 AutoTP 分别处理这四类问题,不能把它们当成同一开关的不同档位。

DeepSpeed 内存与并行方法总览

自绘示意图:四种方法分别改变优化器归属、ZeRO 通信、冻结权重表示或层内切分。

先确定哪个对象超预算

设参数量为 \(N\),数据并行度为 \(P_d\),张量并行度为 \(P_t\)。混合精度 Adam 训练至少要考虑低精度参数、梯度、FP32 主参数以及一阶/二阶矩。忽略临时工作区时,模型状态常数常被粗略写成每参数约 16 bytes;ZeRO 分片和 TP 只改变其中某些对象的归属:

\[ M_{\mathrm{GPU}}\approx M_{\mathrm{param}}+M_{\mathrm{grad}}+M_{\mathrm{optim}}+M_{\mathrm{activation}}+M_{\mathrm{workspace}}. \]

关键问题不是“用了几级 ZeRO”,而是每个对象在 forward、backward 和 optimizer step 的哪个时刻出现、由谁拥有、需要走哪条链路。

ZeRO-Offload:把优化器所有权迁到 CPU

动机与直觉

Adam 的 FP32 主参数和两份矩状态体积很大,optimizer step 也会消耗 GPU 算力。ZeRO-Offload 让 CPU 内存持有优化器状态,并让优化器计算在 CPU 上执行;GPU 继续承担前向和反向。1

这不是“凭空获得显存”,而是一次资源交换:

  • GPU 少放 optimizer state;
  • CPU RAM、CPU 算力和 PCIe/NUMA 路径承担新增压力;
  • pinned memory、分块大小和 CPUAdam 决定传输与更新能否被隐藏。

机制流程

for micro_batch in loader:
    loss = gpu_forward(micro_batch)
    gpu_backward(loss)
    grad_partition = reduce_scatter_gradients()
    async_copy_to_pinned_cpu(grad_partition)

wait_for_required_gradient_partition()
cpu_adam_update(fp32_param_partition, m, v, grad_partition)
async_copy_updated_partition_to_gpu()

数据流是 GPU gradient → pinned CPU buffer → CPUAdam → updated partition → GPU。如果 CPU 更新与 PCIe 传输比 GPU 下一段计算更慢,关键路径只是从显存容量转移成了 host stall。

适用边界

适合 GPU 显存不足、CPU 内存充足并且能控制 NUMA 亲和性的单机或小规模训练。官方教程展示单 GPU 训练 10B GPT-2;这是可行性示例,不是所有单卡都能达到相同吞吐的保证。1

ZeRO++:分别优化三条 ZeRO-3 通信路径

ZeRO++ 不是单一量化开关,而是 qwZ、hpZ、qgZ 三个组件。2

qwZ:量化权重 AllGather

ZeRO-3 前向/反向需要临时 AllGather 参数。qwZ 以 block-based quantization 把 FP16 权重通信为 INT8,接收后反量化:

local FP16 shard
  -> block quantize INT8 + scale
  -> parameter AllGather
  -> dequantize
  -> local layer compute

它降低通信字节,但增加量化 kernel、scale 和临时 buffer。

hpZ:节点内保存次级参数分区

跨节点带宽通常弱于节点内 NVLink/NVSwitch。hpZ 在节点内建立 secondary partition group,让反向参数获取尽量使用节点内副本,以额外显存换掉一次跨节点 AllGather。

qgZ:量化梯度通信

qgZ 在梯度路径上使用量化的 All-to-All/AllGather 组合,降低跨节点 reduce-scatter 等价通信量。三者组合时,论文报告通信量最高降低 4×、吞吐最高提升 2.16×;这两个数是论文网络、模型和规模下的上限,不是通用 SLA。3

MixZ++:让冻结权重持续保持低精度

MixZ++ 面向 LoRA 等“基础权重冻结、少量参数训练”的场景。它继承 qwZ/hpZ,但关键差别是:冻结权重可以一直以低精度形式保存,避免每个使用周期重复量化,也同时降低权重常驻和通信体积。4

frozen base weight:  INT8 shard --AllGather--> INT8 gathered --dequant--> matmul
trainable adapter:   BF16/FP16 local compute -----------------------> update
optimizer state:     only for trainable parameters

官方页面引用的最高 3.3× 来自 Llama-2-70B LoRA、128 张 V100 的评估。它不能外推到全参数训练,也不能证明 INT8 权重路径对所有模型精度无损。

MixZ++ 不是普通 mixed precision 的新名字

普通 BF16/FP16 训练仍可能保留 FP32 主状态;MixZ++ 讨论的是 ZeRO++ 下冻结权重的持续量化布局。先确认训练对象是否主要是 LoRA adapter,再考虑它。

AutoTP:把层规则编译成张量并行

推理与训练是两条路径

早期 AutoTP 教程面向 Hugging Face 推理:识别 Transformer 层并注入列并行/行并行替代层。新的训练 AutoTP 支持 preset、正则 pattern、Hugging Face tp_plan 和自定义 layer spec,并可组合 DP 与 ZeRO 0/1/2;当前官方文档明确不支持 ZeRO Stage 3。56

对象和形状

\(Y=XW\) 为例:

  • 列并行把 \(W\in\mathbb{R}^{d_{in}\times d_{out}}\) 沿 \(d_{out}\) 切开,各 Rank 得到部分输出;
  • 行并行沿 \(d_{in}\) 切开,各 Rank 先算部分和,再做 Reduce/ReduceScatter;
  • Q/K/V、输出投影、MLP gate/up/down 的组合必须保持语义和集合通信配对。
plan = detect_transformer_layers(model)
for layer in plan:
    spec = match_preset_or_regex(layer)
    shard_parameters(layer, spec, tp_group)
    replace_forward(layer, spec.collective_contract)
validate_divisibility_and_tied_weights(model, tp_size)

“自动”只减少规则编写,不消除约束。hidden size、attention heads、KV heads 和 fused parameter layout 仍须可切;自定义层、共享权重和 GQA 可能需要显式 pattern。

四种方法改变的对象所有权与链路

自绘示意图:先定位峰值对象,再改变归属、执行数据流并复验稳态。

如何选择

现象 优先方法 先验证什么
Adam 状态让 GPU OOM ZeRO-Offload CPU RAM、NUMA、PCIe、CPUAdam 是否进入关键路径
ZeRO-3 跨节点通信暴露 ZeRO++ 参数/梯度 collective 分解与节点内外带宽比
LoRA 冻结权重仍占容量和通信 MixZ++ 冻结比例、量化精度、V100/Ampere/Hopper kernel
单层矩阵或单卡算力不足 AutoTP 维度整除、模型规则、collective 和 ZeRO 组合

结论

四种方法改变的是不同对象:ZeRO-Offload 移动优化器,ZeRO++ 改造 ZeRO-3 通信,MixZ++ 固化冻结权重的低精度布局,AutoTP 切分层内张量。正确顺序是先用峰值显存与 trace 找到具体对象,再选择最小机制;不要用一个最高加速数字替代目标集群上的分项测量。


  1. DeepSpeed ZeRO-Offload tutorial,教程文件最早提交于 2020-09-10。 

  2. DeepSpeed ZeRO++ tutorial,教程文件最早提交于 2023-06-23。 

  3. ZeRO++ paper。 

  4. Mixed Precision ZeRO++ tutorial,教程文件最早提交于 2023-08-31。 

  5. Automatic Tensor Parallelism for inference,教程文件最早提交于 2023-02-21。 

  6. Automatic Tensor Parallelism for training。 

评论