ZeRO-1 思想与 Megatron 定位
为什么需要分布式优化器?
在混合精度训练(fp16/bf16)中,模型的前向和反向计算使用低精度参数以节省显存、提高吞吐, 但优化器更新必须在 fp32 精度上进行,否则小梯度会被舍入为零,导致训练不稳定。 这就引入了所谓的 master weight(主参数)机制:优化器在内部维护一套 fp32 副本, 更新完毕后再写回 fp16/bf16 模型参数。
代价是巨大的。以 Adam 为例,对每个参数,优化器需要同时保存:
- fp32 master weight:模型参数的高精度副本(4 bytes/param)
- 一阶矩
m(动量):fp32(4 bytes/param) - 二阶矩
v(方差):fp32(4 bytes/param)
三者合计 12 bytes/param,而模型参数本身只占 2 bytes(bf16)。 对于一个 7B 参数的模型,优化器状态仅 Adam 部分就需要约 84 GB 显存, 远超单张 A100 的 80 GB 容量——更别提还要存放激活值和梯度缓冲区。
传统数据并行的做法与局限
在标准的数据并行(Data Parallel,DP)训练中,每个 rank 保留完整的模型参数和完整的优化器状态。 反向传播结束后,通过一次 All-Reduce 在所有 DP rank 之间同步梯度, 然后各自独立执行相同的优化器 step。
这个方案的问题是:每个 rank 都在做重复的工作,持有重复的状态。 增加 DP 数量并不能减少每卡的显存占用。8 台机器 8×8 = 64 个 rank, 但每张卡的优化器内存压力和单卡时完全相同。
ZeRO-1:分散优化器状态
ZeRO(Zero Redundancy Optimizer)第一阶段(ZeRO-1)的核心思路是: 把优化器状态按 DP rank 切分,每个 rank 只持有并更新 1/DP 的参数分片。 通信方式从 All-Reduce 变成:
- Reduce-Scatter:将各 rank 的梯度汇总,并将结果分散到各 rank, 每个 rank 只得到自己"负责"的那段梯度。
- 各 rank 独立更新自己持有的参数分片(优化器 step)。
- All-Gather:将更新后的参数分片广播给所有 rank,恢复完整参数。
Megatron-LM 的 DistributedOptimizer 正是实现了这一思路。
在 distrib_optimizer.py 的 __init__ 注释中也明确说明:
# distrib_optimizer.py, DistributedOptimizer.__init__ docstring
"""
Distributed optimizer, for all data types (fp16, bf16, and fp32).
The steps in this method create the core mapping between param and grad buffers,
parameters, and parameter shard ranges, that is needed for converting between
model param indexes and main parameter shard indexes. This method also updates
the optimizer parameter groups with the newly created shards.
Args:
per_model_buffers: the implementation of the distributed optimizer is
centered on using a contiguous buffer for communicating grads & params
between the model state and the optimizer state.
data_parallel_group: data-parallel group to use to
all-gather params after optimizer.step().
"""
设模型参数量为 P,数据并行度为 DP,则:
每卡优化器状态 = (fp32_master + m + v) × P × 4 bytes / DP
节省量 = (fp32_master + m + v) × P × 4 bytes × (1 − 1/DP)
其中 fp32_master、m、v 各占 4 bytes/param,三者合计 12 bytes/param。 DP 越大,每卡节省越多,极限情况下(DP → ∞)可接近节省全部 12 bytes/param 的优化器开销。
具体数字:7B 模型,DP=8
假设模型参数量 P = 7×109,数据并行度 DP = 8:
- 传统 DP:每卡优化器状态 ≈ 12 × 7B bytes = 84 GB
- DistributedOptimizer (ZeRO-1):每卡优化器状态 ≈ 84 / 8 = 10.5 GB
- 节省 73.5 GB(节省率 87.5%)
这意味着单张 A100 (80 GB) 原本无法装下的优化器状态,在 8-way DP 的分布式优化器下 每卡只需 10.5 GB,显存空间可以重新分配给更大的批量或更长的序列。
DistributedOptimizer 仅实现了 ZeRO-1(仅分散优化器状态),
不做梯度分散(ZeRO-2)或参数分散(ZeRO-3)。Megatron-LM 已通过张量并行(TP)和流水线并行(PP)
在正交维度上切分了参数,因此无需完整的 ZeRO-3,ZeRO-1 已足够大幅降低显存压力。
这也是 DistributedOptimizer 在模型并行训练中的精确定位。
类继承结构
三层继承体系
Megatron-LM 的优化器体系由三个核心类构成,形成严格的继承链:
MegatronOptimizer → MixedPrecisionOptimizer → DistributedOptimizer。
每一层各司其职,职责边界清晰。
第一层:MegatronOptimizer(抽象基类)
定义于 optimizer.py 第 108 行,是纯抽象基类(继承自 ABC)。
它封装了一个底层 torch.optim.Optimizer 实例,并通过 property 将
state、param_groups 代理到内部优化器,
使外层代码可以像访问标准 PyTorch 优化器一样访问 Megatron 优化器。
# optimizer.py, line 108
class MegatronOptimizer(ABC):
"""
Base class for all Megatron optimizers.
Args:
optimizer (torch.optim.Optimizer): base optimizer such as Adam or SGD.
config (OptimizerConfig): configuration object for optimizer.
init_state_fn (Callable, optional): function to initialize state in the optimizer.
"""
def __init__(
self,
optimizer: torch.optim.Optimizer,
config: OptimizerConfig,
init_state_fn: Callable = lambda x: None,
):
self.optimizer = optimizer
self.config = config
self.init_state_fn = init_state_fn
该层定义的抽象方法包括:prepare_grads()、step_with_ready_grads()、
zero_grad()、get_loss_scale()、reload_model_params()、
state_dict()、load_state_dict()、step()、
sharded_state_dict()。所有子类都必须实现这些接口。
此外还提供了梯度裁剪(clip_grad_norm)、梯度范数计算(get_grad_norm)、
零梯度计数(count_zeros)等通用工具方法,子类无需重复实现。
第二层:MixedPrecisionOptimizer(混合精度公共逻辑)
定义于 optimizer.py 第 453 行。它是 DistributedOptimizer
和 Float16OptimizerWithFloat16Params 的共同父类,
两者的差别在于参数分片方式——前者按 DP rank 切分,后者持有完整参数。
# optimizer.py, line 453
class MixedPrecisionOptimizer(MegatronOptimizer):
"""Base class for both the float-16 and the distributed optimizer.
Args:
optimizer (torch.optim.Optimizer): base optimizer such as Adam or SGD.
config (OptimizerConfig): configuration object for optimizer.
grad_scaler (MegatronGradScaler): used for scaling gradients. Note that
this can be None. This case happens when `bf16 = True` and we don't
use any loss scale. Note that for `bf16 = True`, we can have
a constant gradient scaler. Also for `bf16 = False`, we
always require a grad scaler.
init_state_fn (Callable, optional): function to initialize state in the optimizer.
"""
该层的核心贡献是实现了完整的 step() 流程骨架:
- 调用
prepare_grads():将模型梯度拷贝到 fp32 主梯度,并检测 inf/NaN; - 调用
clip_grad_norm()(若开启); - 调用
step_with_ready_grads():执行内层优化器 step,将主参数写回模型参数。
同时维护 grad_scaler(梯度缩放器,fp16 训练必须;bf16 可为 None)
和 found_inf(溢出检测张量)。
子类需要实现三个 hook:_copy_model_grads_to_main_grads()、
_collect_main_grad_data_for_unscaling()、
_copy_main_params_to_model_params(),
分别负责梯度拷贝、unscale 数据收集和参数写回,这正是分布式与非分布式实现的差异所在。
第三层:DistributedOptimizer(ZeRO-1 分布式优化器)
定义于 distrib_optimizer.py 第 95 行,这是整个体系的核心实现层。
# distrib_optimizer.py, line 95
class DistributedOptimizer(MixedPrecisionOptimizer):
"""Distributed optimizer, for all data types (fp16, bf16, and fp32).
See __init__() below for argument details.
"""
# enumerates fully reshardable optimizer formats (as opposed to formats
# which depend on the internal optimizer buffers structure)
checkpoint_fully_reshardable_formats: set[str] = {
'fully_reshardable',
'fully_sharded_model_space',
'fsdp_dtensor',
}
该层在 MixedPrecisionOptimizer 的骨架之上,补充了 ZeRO-1 所需的全部分布式逻辑:
-
参数范围映射(Range Mapping):通过一系列
@classmethod(_build_gbuf_range_map、_build_model_gbuf_param_range_map等) 建立 grad buffer 分片 与 参数 之间的精确映射, 为 Reduce-Scatter 和 All-Gather 提供地址索引。 -
主参数分片(Shard Main Params):通过
_build_model_and_main_param_groups()为每个 rank 只创建其"负责"范围内的 fp32 主参数切片(shard_fp32_from_float16_groups), 而非复制完整参数。 -
梯度通信覆盖:重写
_copy_model_grads_to_main_grads()等方法, 操作对象从完整参数变成 grad buffer 的分片视图。 -
梯度统计范围扩展:重写
get_grad_stats_parallel_group(), 返回整个分布式优化器实例的通信组(而非仅模型并行组),确保梯度范数在所有 DP rank 间全局归一化。
DistributedOptimizer 仅支持 Adam
在 __init__ 中有明确断言(第 524-530 行):
assert (
isinstance(optimizer, (Adam, torch.optim.AdamW, HybridDeviceOptimizer))
or optimizer is None
), (
"Only Adam and HybridDeviceOptimizer currently supported, "
"due to checkpointing requirements."
)
这是因为分布式优化器的 checkpoint 序列化逻辑高度依赖 Adam 优化器状态的具体结构(exp_avg、
exp_avg_sq),切换为其他优化器时需重新实现 sharded_state_dict()。
ParamAndGradBuffer 与 Bucket 划分
Buffer 的整体布局
_ParamAndGradBuffer 最核心的设计思想是:把同一数据类型的所有参数/梯度平铺到一块连续的 GPU 内存,
再在这块内存上按固定大小切出若干 bucket。NCCL 的 AllReduce / ReduceScatter 操作直接接收一个连续的指针和长度,
因此这种布局可以让集合通信零拷贝地覆盖整块 grad buffer,避免了把散落各处的梯度张量逐个收集再通信的开销。
具体地,param_data(参数存储)和 grad_data(梯度存储)各自是一个独立的
torch.Tensor,元素数量均为 self.numel。每个参数通过
param_index_map[param] = (start, end, bucket_id) 记录其在这块大 tensor 中的起止偏移,
从而让参数 tensor 成为该区间的一个 view,而非独立的内存分配。
Bucket 划分算法:逆序遍历 + 尺寸触发
划分过程发生在 _ParamAndGradBuffer.__init__ 中,核心逻辑如下:
for param in params[::-1]:
# 逆序遍历——大致与反向传播顺序一致
this_numel = param.data.nelement()
param_start_index = _pad_start_of_param_if_needed(param_start_index)
# 若当前参数需要独占 bucket(如 shared_embedding),先封闭上一个 bucket
if _does_param_require_new_bucket(param) and len(bucket_params) > 0:
param_start_index = _update_bucket_metadata(param_start_index)
param_end_index = param_start_index + this_numel
self.param_index_map[param] = (param_start_index, param_end_index, bucket_id)
bucket_params.add(param)
# 超过 bucket_size 阈值 → 封闭当前 bucket,开启下一个
if (
bucket_size is not None
and (param_end_index - bucket_start_index) >= bucket_size
) or _does_param_require_new_bucket(param):
bucket_end_index = _update_bucket_metadata(param_end_index)
param_start_index = bucket_end_index
else:
param_start_index = param_end_index
# 最后一组参数形成尾 bucket
if len(bucket_params) > 0:
bucket_end_index = _update_bucket_metadata(param_end_index)
两级 Padding 机制
代码中存在两处对齐填充,各有不同目的:
-
参数起始对齐(
_pad_start_of_param_if_needed):每个参数在 buffer 中的起始偏移必须是 64 的倍数(对应 128-byte 地址对齐,因为 ≥ 16-bit 精度每元素 2 字节)。这是给 DistributedOptimizer 使用的,普通 DDP 不需要。 -
Bucket 末尾对齐(
_pad_end_of_bucket_if_needed):DistributedOptimizer 要求每个 bucket 的元素数能被DP_world_size整除,以便均等分片。基准 divisor 为lcm(dp_world_size, 128);若开启pad_buckets_for_high_nccl_busbw,则进一步扩大到lcm(dp_world_size, 128, 2^16),确保 NCCL ring 算法中每个 rank 的消息块是 2 的幂次,以获得最高总线带宽。
def _pad_end_of_bucket_if_needed(bucket_end_index: int) -> int:
if self.ddp_config.use_distributed_optimizer:
if self.ddp_config.pad_buckets_for_high_nccl_busbw:
# NCCL ring 算法:消息大小 = bucket_size / dp_size
# 需要整除 2^16 才能获得高总线带宽
bucket_size_divisor = math.lcm(self.data_parallel_world_size, 128, 2**16)
else:
bucket_size_divisor = math.lcm(self.data_parallel_world_size, 128)
return _pad(bucket_end_index, bucket_size_divisor)
return bucket_end_index
def _pad_start_of_param_if_needed(param_start_index: int) -> int:
if self.ddp_config.use_distributed_optimizer:
# 128-byte 对齐(64 个 >=16-bit 元素)
return _pad(param_start_index, 64)
return param_start_index
Buffer → Buckets → Params 层次结构图
反向传播遵循 LIFO(后进先出) 顺序:越靠近输出层(模型列表末尾)的参数,其梯度越先计算完成。
params[::-1] 的逆序遍历使得最早产生梯度的参数被放入第一个 bucket(索引 0),
从而当该 bucket 被填满时,其中的梯度几乎全部已就绪,可以立刻触发 AllReduce / ReduceScatter,
与仍在计算中的前层梯度流水线并行,显著缩短端到端通信等待时间。
最终内存大小断言
所有 bucket 划分完成后,代码做如下验证:
self.numel = bucket_end_index # 含 padding 的总元素数
self.numel_unpadded = sum(per_bucket_numel_unpadded) # 不含 padding
assert self.numel_unpadded <= self.numel
if self.ddp_config.use_distributed_optimizer:
# 关键断言:整个 grad buffer 必须能被 DP size 整除
# 每个 bucket 已对齐,因此全局也自然对齐
assert self.numel % self.data_parallel_world_size == 0
else:
assert self.numel == self.numel_unpadded # 非 dist opt 时无 padding
参数分片——每个 DP rank 持有什么
Bucket 级别的均等切分
DistributedOptimizer 的核心思路是:对每个 bucket 的 grad buffer 做 ReduceScatter,
每个 DP rank 只接收并持有 1/DP_size 的梯度片段。
正是为了让这个切分均匀,前一节才对 bucket 进行对齐 padding。
切分在 _build_model_gbuf_range 中完成:
@classmethod
def _build_model_gbuf_range(cls, param_and_grad_buffer, bucket_index):
data_parallel_rank = param_and_grad_buffer.data_parallel_group.rank()
data_parallel_world_size = param_and_grad_buffer.data_parallel_group.size()
bucket = param_and_grad_buffer.buckets[bucket_index]
gbuf_size = bucket.grad_data.numel()
assert gbuf_size % data_parallel_world_size == 0
max_gbuf_range_size = gbuf_size // data_parallel_world_size # 每 rank 的片段大小
# 计算所有 DP rank 的世界坐标区间
gbuf_world_all_ranges = []
for r in range(data_parallel_world_size):
gbuf_world_start = r * max_gbuf_range_size
gbuf_world_end = min(gbuf_size, gbuf_world_start + max_gbuf_range_size)
# 加上 bucket 在整个 grad_buffer 中的偏移
gbuf_world_range = Range(
gbuf_world_start + bucket.offset, gbuf_world_end + bucket.offset
)
gbuf_world_all_ranges.append(gbuf_world_range)
# 当前 rank 持有的区间
gbuf_world_range = gbuf_world_all_ranges[data_parallel_rank]
Range 类:四种视角的统一抽象
Range 是 distrib_optimizer.py 中定义的轻量数据类,仅持有 start、end、size 三个字段,
并提供 normalize(start) 方法用于坐标系平移:
class Range:
"""表示从完整张量中索引一个分片的起止区间。"""
def __init__(self, start: int, end: int):
self.start = start
self.end = end
self.size = end - start # 等价于 __len__
def normalize(self, start: int = 0):
"""将区间平移,使新的 start 对齐到指定偏移。"""
return Range(start, start + self.size)
def __str__(self):
return "%d,%d [%d]" % (self.start, self.end, self.size)
def __len__(self):
return self.end - self.start
对于每个落在本 rank 分片内(或横跨边界)的参数,_build_model_gbuf_param_range_map
会建立一个包含四种坐标视角的字典:
param_range_map[param] = {
"gbuf_world": param_world_range, # 视角①
"gbuf_world_in_bucket": param_world_range_in_bucket, # 视角②
"gbuf_local": param_local_range, # 视角③
"param": sub_param_range, # 视角④
}
四种视角详解
grad_data tensor 创建参数梯度 view,以及 copy_main_to_model_params 时的索引。
bucket.offset,转换为在当前 bucket grad buffer 内的相对偏移。
AllGather / ReduceScatter 操作直接在单个 bucket 的 tensor 上执行,因此需要此视角。
sub_param_range 精确描述本 rank 持有的那一段在参数内的位置,
用于 model → main param 的拷贝和 AllGather 后写回参数。
4 个 DP rank 切割 bucket 的示意图
| 字段 | rank 0(持有 [32,36)) | rank 1(持有 [36,40)) |
|---|---|---|
| ① gbuf_world | Range(34, 36) [裁剪到本 rank] | Range(36, 39) [裁剪到本 rank] |
| ② gbuf_world_in_bucket | Range(2, 4) [减去 bucket.offset=32] | Range(4, 7) [减去 bucket.offset=32] |
| ③ gbuf_local | Range(2, 4) [在本 rank 4 元素片中] | Range(0, 3) [在本 rank 4 元素片中] |
| ④ param(P 共 5 元素) | Range(0, 2) [P 的前 2 个元素] | Range(2, 5) [P 的后 3 个元素] |
rank 0 只拿到 P 的前 2 个元素,rank 1 只拿到 P 的后 3 个元素,两者合并才是完整的 P 梯度。
计算四种视角的源码
实际计算逻辑在 _build_model_gbuf_param_range_map 中,逐行解读:
for param, param_world_indexes in param_world_index_map.items():
param_world_start, param_world_end, _ = param_world_indexes
# 计算参数与本 rank 分片的交集(local 坐标,从 0 开始)
param_local_start = max(0, param_world_start - gbuf_world_range.start)
param_local_end = min(gbuf_world_range.size, param_world_end - gbuf_world_range.start)
# 只处理有交集的参数(完全在其他 rank 分片内的参数直接跳过)
if param_local_end > param_local_start:
# ③ gbuf_local:在本 rank 持有的小 tensor 中的偏移(从 0 起算)
param_local_range = Range(param_local_start, param_local_end)
# ① gbuf_world:平移回全局 grad_data 坐标
param_world_range = param_local_range.normalize(
param_local_start + gbuf_world_range.start
)
# ② gbuf_world_in_bucket:减去 bucket 的起始偏移
param_world_range_in_bucket = Range(
param_world_range.start - bucket_offset,
param_world_range.end - bucket_offset,
)
# ④ param:在参数自身 tensor 内的起始位置(参数可能被切成两半)
sub_param_start = max(0, gbuf_world_range.start - param_world_start)
sub_param_range = param_local_range.normalize(sub_param_start)
param_range_map[param] = {
"gbuf_world": param_world_range,
"gbuf_world_in_bucket": param_world_range_in_bucket,
"gbuf_local": param_local_range,
"param": sub_param_range,
}
DistributedOptimizer 故意不按参数边界切分,而是按 bucket 的 1/DP_size 位置硬切。
这样可以保证每个 rank 的工作量完全均等(各持有相同数量的梯度元素,优化器更新步骤负载均衡),
代价是一个参数可能被两个相邻 rank 分别持有不同子区间,需要用上面四种 Range 做精确映射。
这也是 normalize() 方法存在的核心原因——各坐标系之间的转换本质上是一次偏移量的平移。
混合精度主参数分片
DistributedOptimizer 的核心职责之一,是在每个 DP rank 上只为自己
负责更新的那一小片参数维护 fp32 副本(主参数,main param)。
这一小片由 Reduce-Scatter 决定:bucket 大小为 N 个元素,
DP world size 为 D,则本 rank 恰好负责 N/D 个元素。
fp32 主参数不存储在连续 buffer 中——每个 rank 只需要一小片,
独立 clone 成 float32 张量即可;而 param_data / grad_data
这两个连续 buffer 是所有 rank 共享的通信介质,不需要也不应该膨胀成 fp32。
五组数据结构
_build_model_and_main_param_groups()(第 305–467 行)
在初始化阶段填充以下五个列表,每个列表按 optimizer param group 分组:
| 变量名 | 存储内容 | dtype | 用途 |
|---|---|---|---|
| model_float16_groups | 完整的 fp16/bf16 模型参数(原始对象) | bf16 / fp16 | 读取 main_grad,写回更新后的参数(通过 shard) |
| model_fp32_groups | 完整的 fp32 模型参数(原始对象) | fp32 | 纯 fp32 训练场景;shard 直接是其视图 |
| shard_float16_groups | bf16 参数展平后的本 rank 分片视图 | bf16 / fp16 | precision-aware optimizer 模式下直接传入优化器 |
| shard_fp32_groups | fp32 参数展平后的本 rank 分片视图 | fp32 | 优化器直接更新这个视图(就地写回模型参数) |
| shard_fp32_from_float16_groups | 从 bf16 参数对应的 fp32 主参数副本(clone) | fp32 | Adam 实际更新的目标;更新后写回 bf16 buffer |
分片切片逻辑
每个参数对应的本 rank 分片范围由
gbuf_range["param_map"][model_param]["param"]
给出(一个 Range 对象,含 .start 和 .end)。
切片与 clone 过程如下:
# distrib_optimizer.py L359-L415(bf16 / fp16 参数路径,已删去 FP8 分支)
if model_param.type() in ['torch.cuda.HalfTensor', 'torch.cuda.BFloat16Tensor']:
# 1. 生成 shard_model_param:bf16 展平视图,截取本 rank 负责的区间
shard_model_param = model_param.detach().view(-1)[
param_range.start : param_range.end
]
tensor_parallel.copy_tensor_model_parallel_attributes(
shard_model_param, model_param
)
# 2. 生成 shard_main_param:clone 并升精度到 fp32(真正的"主参数")
shard_main_param = shard_model_param.clone().float()
tensor_parallel.copy_tensor_model_parallel_attributes(
shard_main_param, model_param
)
# 3. 挂载标记属性到原始 bf16 参数对象
model_param.main_param = shard_main_param
model_param.main_param_sharded = True
# 4. 放入对应 group 列表
model_float16_params_this_group.append(model_param)
shard_float16_params_this_group.append(shard_model_param)
shard_fp32_from_float16_params_this_group.append(shard_main_param)
其中 param_range 是该参数在本 rank 所拥有的 gbuf 本地视图内的子范围——
即在 gbuf_local(大小为 bucket_numel / DP_size)内的偏移。
由于 grad buffer 分区不尊重参数边界,一个参数可能只有一部分落在本 rank 的区间里,
param_range 正好记录了这个参数中属于本 rank 的子区间。
-
model_param.main_param:指向本 rank 对应的 fp32 主参数分片 (shard_main_param)。初始化时先在所有 bf16 参数上置为None(第 593 行),待_build_model_and_main_param_groups()执行后, 有优化器状态的参数会将其覆盖为真实张量;没有优化器状态的参数(其他 rank 负责的参数) 保持None。grad norm 计算时遍历main_param而不是 bf16 参数本身,从而得到 fp32 精度的 norm。 -
model_param.main_param_sharded = True:标志该参数的主参数是 分片的(只有本 rank 持有自己的那一片),区别于非分布式优化器中 每个 rank 持有完整副本的情形。检查点保存 / 加载逻辑通过此标志 判断是否需要做 DP 维度的 gather。
fp32 主参数为何不存于连续 buffer
param_data 和 grad_data 这两个连续 buffer 以 bf16
存储(或 fp32,视配置),所有 DP rank 共享同一套布局,是 All-Gather / Reduce-Scatter
通信的载体。如果把 fp32 主参数也放进连续 buffer,则:
- 每个 rank 需要存一份完整 fp32 buffer(内存 × DP_size × 2 倍精度膨胀);
- All-Gather / Reduce-Scatter 通信量翻倍;
- Adam 的 exp_avg / exp_avg_sq 也将跟着膨胀,最终 ZeRO-1 的内存节省完全失去意义。
实际上每个 rank 只需要 1/DP_size 的 fp32 主参数,因此直接
.clone().float() 产生一个独立的小张量即可,不必对齐到任何连续 buffer。
优化器(Adam)在这些小张量上原地更新,更新完成后再按原路写回 bf16 param buffer,
随后一次 All-Gather 恢复完整参数。
完整训练步骤追踪
以 DP=4、单个 bucket 含 1024 个 bf16 元素 为例,
追踪从反向传播到下一次前向的完整数据流。
每个 rank 拥有 1024 / 4 = 256 个元素的分片。
param.main_grad
是 grad_data buffer 中的一个视图(view),反向传播时梯度直接就地写入连续内存,
无需额外拷贝。DDP 的 hook(param.grad_added_to_main_grad 标志控制)
确保梯度累加到正确位置。
model_chunk.start_grad_sync() +
finish_grad_sync()
触发(位于 distributed_data_parallel.py 第 581/593 行)。
每个 rank 的 grad_data buffer(1024 bf16)参与 Reduce-Scatter:
对应位置的梯度在 DP=4 个 rank 之间求和,然后分发,
每个 rank 得到属于自己那 256 元素的已归约梯度。
grad[0..1023]
grad[0..1023]
grad[0..1023]
grad[0..1023]
grad[0..255]
sum of 4
grad[256..511]
sum of 4
grad[512..767]
sum of 4
grad[768..1023]
sum of 4
_copy_model_grads_to_main_grads()
由 MixedPrecisionOptimizer.prepare_grads()(optimizer.py 第 550 行)调用。
对每个 bf16 参数,从 model_param.main_grad(grad_data buffer 的视图)
取本 rank 分片,升精度后赋给 fp32 主参数的 .grad:
# distrib_optimizer.py L2428-L2442
param_range_map = self._get_model_param_range_map(model_param)
param_range = param_range_map["param"] # 参数内子区间
model_grad = model_param.main_grad # bf16 view into grad_data
shard_model_grad = model_grad.view(-1)[
param_range.start : param_range.end # 本 rank 对应的 256 个梯度元素
]
shard_main_param.grad = shard_model_grad.float() # bf16 -> fp32
clip_grad_norm()
计算全局 grad norm(跨 TP/DP rank all-reduce),
随后 Adam.step() 在本 rank 的 256 个 fp32 主参数上执行更新。
各 rank 完全独立执行,不涉及任何通信。
_copy_main_params_to_model_params()
由 MixedPrecisionOptimizer.step_with_ready_grads()(optimizer.py 第 601 行)
在 optimizer.step() 之后立即调用。
# distrib_optimizer.py L2483-L2498
world_range = param_range_map["gbuf_world_in_bucket"] # 全局 bucket 视图内偏移
model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data
shard_model_param = model_param_buffer.view(-1)[
world_range.start : world_range.end # param_data 中本 rank 负责的区间
]
shard_model_param.data.copy_(shard_main_param) # fp32 -> bf16(自动截断精度)
param_range(参数内子区间),
参数写回(阶段 5)使用 gbuf_world_in_bucket(全局 bucket 内绝对偏移)。
两套坐标系由 _build_model_gbuf_param_range_map() 统一预计算。
model_chunk.start_param_sync()
由 DistributedOptimizer.step_with_ready_grads()(第 2662 行)调用。
每个 rank 贡献自己更新好的 256 个 bf16 元素,All-Gather 后所有 rank 的
param_data buffer 都恢复为完整的 1024 bf16 元素。
模型参数(model_param.data)是 param_data buffer 的视图,
因此自动得到更新,无需额外拷贝。
param[0..255]
already updated
param[256..511]
already updated
param[512..767]
already updated
param[768..1023]
already updated
param[0..1023] bf16
param[0..1023] bf16
param[0..1023] bf16
param[0..1023] bf16
overlap_param_gather=True,All-Gather 延迟到下一次前向的
pre-hook 中发起,与计算重叠,进一步隐藏通信延迟。
grad_data buffer 的第 [rank × 256, (rank+1) × 256) 位置,
在 Reduce-Scatter 之后存放的就是该区间的已归约梯度。
_copy_model_grads_to_main_grads() 中用的 param_range
坐标系(参数内子区间)是相对于 gbuf_local 视图计算的,
两者的边界由 _build_model_gbuf_param_range_map() 预先对齐,
因此直接按 param_range.start:param_range.end 切片就能拿到正确的已归约梯度,
无需额外的偏移换算。注意这里是 param.main_grad.view(-1) 的切片,
而非整个 grad_data buffer 的切片,坐标系从参数在 grad_data 中的起点开始。
上述 6 个阶段以同步方式描述。实际训练中两个 overlap 开关可以把通信隐藏在计算后面:
-
overlap_grad_reduce=True:Reduce-Scatter(阶段 2)在反向传播的 bucket 填满时 立即异步发起,与后续层的反向计算重叠。bucket 越多重叠效果越好。 -
overlap_param_gather=True:All-Gather(阶段 6)延迟到下一次迭代的前向 pre-hook 中发起,与第一层前向计算重叠。此时step_with_ready_grads()中不调用start_param_sync(),转而由 zero_grad() / forward pre-hook 负责。
这两个开关不改变数据正确性,只改变通信与计算的调度时序, 是大规模训练的重要吞吐量优化手段。
bf16 训练时 grad_scaler 可以为 None(无需动态 loss scaling)。
fp16 训练时 prepare_grads() 在拷贝梯度之后还会调用
_unscale_main_grads_and_check_for_nan(),
对 fp32 主梯度做反缩放并检测 inf/nan,若发现则跳过本次优化器步骤并更新 loss scale。
这一检测通过 all-reduce found_inf 张量跨 MP rank 同步,
保证所有 rank 在是否跳过步骤上达成一致。
Overlap 与异步优化
DistributedOptimizer 提供了两个关键的 overlap 选项,可以大幅隐藏通信延迟:
overlap_grad_reduce(梯度 reduce-scatter 与反向传播 overlap)
和 overlap_param_gather(参数 all-gather 与前向传播 overlap)。
两者配合使用时,整个通信开销几乎可以完全隐藏在计算之后。
7.1 overlap_grad_reduce:Bucket 级别异步 reduce-scatter
在没有 overlap 时,反向传播结束后必须等所有参数的梯度都就绪,才能发起一次全局的 reduce-scatter。
而 overlap_grad_reduce=True 的核心思想是:将参数分成若干 bucket,
每当一个 bucket 内的所有参数梯度计算完毕,立即异步触发该 bucket 的 reduce-scatter,
不等待其他 bucket 完成。这样后续 bucket 的梯度计算可以与前面 bucket 的通信并行进行。
触发的入口是注册在每个参数梯度累加器(AccumulateGrad)上的 backward post-hook:
# megatron/core/distributed/distributed_data_parallel.py
def _make_backward_post_hook(self, param: torch.nn.Parameter):
def hook(*unused):
if param in self.param_to_bucket_group:
bucket_group = self.param_to_bucket_group[param]
if bucket_group.ddp_config.overlap_grad_reduce:
assert param.grad is not None, \
'param.grad being None is not safe when overlap_grad_reduce is True'
if param.grad is not None and (
not param.grad_added_to_main_grad or getattr(param, 'zero_out_wgrad', False)
):
param.grad.data.record_stream(torch.cuda.current_stream())
param.main_grad.add_(param.grad.data)
param.grad = None
if bucket_group.ddp_config.overlap_grad_reduce:
bucket_group.register_grad_ready(param) # 触发时机
return hook
register_grad_ready() 维护一个 params_with_grad 集合,
当集合大小等于该 bucket group 的总参数数时,立即调用 start_grad_sync():
# megatron/core/distributed/param_and_grad_buffer.py
def register_grad_ready(self, param: torch.nn.Parameter):
assert self.ddp_config.overlap_grad_reduce
if self.is_last_microbatch: # 仅最后一个 microbatch 才同步
self.params_with_grad.add(param)
# 当 bucket 内所有参数梯度就绪 -> 立即触发通信
if len(self.params_with_grad) == len(self.params):
self.start_grad_sync()
start_grad_sync() 内部使用 _coalescing_manager 将同一
bucket group 内多个 bucket 的通信内核合并,然后异步发起 reduce-scatter(DistributedOptimizer)
或 all-reduce(非分布式优化器):
# megatron/core/distributed/param_and_grad_buffer.py(start_grad_sync 核心)
async_op = (
self.ddp_config.overlap_grad_reduce
and self.ddp_config.num_distributed_optimizer_instances == 1
)
with _coalescing_manager(communication_group, async_ops=async_op) as cm:
for idx, bucket in enumerate(self.buckets):
if self.ddp_config.use_distributed_optimizer:
local_data_view = self.cached_grad_buffer_shard_list[idx][
self.intra_distributed_optimizer_instance_rank
]
grad_reduce_handle = dist_reduce_scatter_func(
local_data_view, # output: 本 rank 的梯度分片
bucket.grad_data, # input: 完整梯度 buffer
op=reduce_op,
group=communication_group,
async_op=async_op, # True -> 立即返回 handle,不阻塞
)
7.2 overlap_param_gather:all-gather 与前向传播 overlap
优化器 step 完成后,每个 rank 只持有自己分片的 fp32 参数,需要通过 all-gather
重建完整 bf16 参数才能进行下一次前向传播。overlap_param_gather=True
将这一步拆分为:
- step 结束时,异步发起第一个 bucket 的 all-gather(而非同步等待完成);
- 前向传播中,每个子模块执行前,forward pre-hook 等待该模块所需参数所在 bucket 的 all-gather 完成,同时异步发起下一个 bucket 的 all-gather;
- 这样 all-gather 与前向计算流水线交替进行。
在 step_with_ready_grads() 中,根据 overlap 标志决定是否同步:
# megatron/core/optimizer/distrib_optimizer.py
@torch.no_grad()
def step_with_ready_grads(self) -> bool:
"""Step the optimizer with ready gradients, return successful.
Under the hood, either launch synchronous param all-gathers or get ready to launch
asynchorous all-gathers that get overlapped with the next forward pass.
"""
update_successful = super().step_with_ready_grads()
# ...
# 如果不开 overlap,在这里做同步 all-gather
# 如果开 overlap,第一个 all-gather 将在下一个 optimizer.zero_grad() 中
# 以异步方式发起,后续由 forward pre-hook 驱动流水线
if not self.ddp_config.overlap_param_gather:
for model_chunk in self.model_chunks:
model_chunk.start_param_sync() # 同步,阻塞直到完成
BucketGroup.start_param_sync() 根据 overlap_param_gather
决定 async_op,将分片参数通过 all-gather 写回完整参数 buffer:
# megatron/core/distributed/param_and_grad_buffer.py
def start_param_sync(self, force_sync: bool = False):
assert self.ddp_config.use_distributed_optimizer
async_op = self.ddp_config.overlap_param_gather and not force_sync
with _coalescing_manager(
self.intra_distributed_optimizer_instance_group, async_ops=async_op
) as cm:
for idx, bucket in enumerate(self.buckets):
local_data_view = self.cached_param_buffer_shard_list[idx][
self.intra_distributed_optimizer_instance_rank
]
dist_all_gather_func(
bucket.param_data, # output: 重建后的完整参数 buffer
local_data_view, # input: 本 rank 持有的参数分片
group=self.intra_distributed_optimizer_instance_group,
async_op=async_op,
)
if async_op:
self.param_gather_handle = cm # 保存 handle,供 finish_param_sync 等待
else:
self.param_gather_handle = None
self.param_gather_dispatched = True
forward pre-hook 在每个模块执行前调用 finish_param_sync(),
等待当前 bucket 完成并提前发起下一个 bucket:
# megatron/core/distributed/distributed_data_parallel.py
def _make_forward_pre_hook(self):
def hook(module, *unused):
for param in module.parameters(recurse=False):
if param not in self.param_to_bucket_group:
continue
bucket_group = self.param_to_bucket_group[param]
skip_next_bucket_dispatch = (
self.ddp_config.align_param_gather
or self.overlap_param_gather_with_optimizer_step
)
# 等待本 bucket 的 all-gather 完成,同时触发下一个 bucket 的异步 all-gather
self.param_to_bucket_group[param].finish_param_sync(
skip_next_bucket_dispatch=skip_next_bucket_dispatch
)
return hook
# finish_param_sync 核心逻辑(param_and_grad_buffer.py):
def finish_param_sync(self, skip_next_bucket_dispatch: bool = False):
if not self.param_gather_dispatched:
self.start_param_sync() # 如果尚未发起,先发起
if self.param_gather_handle is not None:
self.param_gather_handle.wait() # 等待完成
self.param_gather_handle = None
# 立即触发下一个 bucket 的异步 all-gather,实现流水线
if self.next_param_gather_bucket_group is not None \
and not skip_next_bucket_dispatch:
self.next_param_gather_bucket_group.start_param_sync()
7.3 时序对比图
以下时序图展示无 overlap 与有 overlap(同时开启两个选项)时,单次训练迭代的时间线对比:
等所有梯度就绪
阻塞等待完成
更新全量 fp32 分片
阻塞等待完成
计算梯度
计算梯度
计算梯度
异步进行
异步进行
异步进行
等所有 RS handle 完成后执行
wait AG[0]
wait AG[1]
wait AG[N]
异步进行
异步进行
异步进行
1. handle 必须在使用参数前等待完成。
finish_param_sync() 负责调用 param_gather_handle.wait()。
若跳过此步骤(如在 align_param_gather 或
overlap_param_gather_with_optimizer_step 模式下提前跳过 dispatch),
参数 buffer 中的数据可能仍是上一个迭代的旧值,导致前向计算结果错误。
2. 参数注册顺序须与前向执行顺序一致。
源码中有明确警告:若 next bucket 的 AG 已经被 dispatch 过
(param_gather_dispatched == True),说明参数注册顺序与实际前向顺序
不匹配,会导致 overlap 性能退化甚至数据不一致:
"The next bucket's parameter all-gather operation has already been dispatched.
This may be caused by a mismatch between the order of parameter registration
and forward pass execution, which will hurt the communication-computation overlap performance."
3. 梯度 overlap 仅在最后一个 microbatch 触发。
register_grad_ready() 检查 is_last_microbatch,
梯度累积阶段(非最后 microbatch)不会提前触发 reduce-scatter,
避免中间梯度被错误通信。
4. 多 DistOpt 实例场景下需额外 stream 同步。
当 num_distributed_optimizer_instances > 1 时,
reduce-scatter 在单独的 communication stream 上执行,需等待默认 stream
完成梯度计算后才能发起(communication_stream.wait_stream(torch.cuda.default_stream())),
且仅该场景下 async_op 仍为 True(通过 stream 隔离实现异步)。
对比与总结
8.1 内存占用对比
设模型参数量为 P,数据并行度为 DP。
Float16OptimizerWithFloat16Params(即非分布式优化器)在每个 rank 上
均持有完整的 fp32 主参数及 Adam 状态;DistributedOptimizer 将这三份
fp32 数据分片到各 DP rank,每个 rank 只持有 P/DP 的份额。
前者在 fp32_from_float16_groups 中为每个 bf16 参数存储了完整的 fp32 副本,
后者则在 shard_fp32_from_float16_groups 中只存储本 rank 负责的分片。
| 组件 | Float16Optimizer(每 rank) | DistributedOptimizer(每 rank) |
|---|---|---|
| 模型参数(bf16) | P × 2 bytes(全量,常驻) | P × 2 bytes(all-gather 后临时重建,常驻于 param buffer) |
| fp32 主参数 fp32_from_float16_groups |
P × 4 bytes(全量) | P/DP × 4 bytes(本 rank 分片) |
Adam exp_avg(fp32) |
P × 4 bytes(全量) | P/DP × 4 bytes(本 rank 分片) |
Adam exp_avg_sq(fp32) |
P × 4 bytes(全量) | P/DP × 4 bytes(本 rank 分片) |
| 梯度 buffer(bf16) | P × 2 bytes(全量) | P × 2 bytes(全量;RS 后本 rank 的梯度分片留在 grad_data 内) |
| fp32 状态合计(可节省部分) | P × 12 bytes(baseline) |
P × 12 / DP bytes 节省约 (1 − 1/DP) × P × 12 bytes |
注:bf16 参数(2 bytes × P)在两种优化器中均以全量存在(前者常驻,后者通过 all-gather 重建),
因此 bf16 参数部分无节省。可节省的是 fp32 主参数 + Adam m/v 合计 12 bytes/元素的
(1 − 1/DP) 比例。以 7B 模型、DP=8 为例:fp32 状态共约 84 GB,
DistributedOptimizer 将其压缩至约 10.5 GB/rank,节省约 73.5 GB。
8.2 通信量对比
两种方案的总通信量在数学上是等价的(Ring All-Reduce = Reduce-Scatter + All-Gather), 区别在于时序灵活性和峰值内存行为:
| 维度 | All-Reduce(非分布式优化器) | Reduce-Scatter + All-Gather(DistributedOptimizer) |
|---|---|---|
| 总通信量(per rank) | 2 × (N−1)/N × P |
RS:(N−1)/N × P AG:(N−1)/N × P 合计:2 × (N−1)/N × P(完全相同) |
| 通信时机 | 等所有梯度就绪后一次性发起(低灵活性) | 按 bucket 流水线触发(高灵活性,可 overlap) |
| 梯度 buffer 峰值 | 需全量梯度 buffer(P × 2 bytes),all-reduce 期间持续占用 | 同样需全量(RS 前),RS 完成后各 bucket 的梯度分片已写回 grad_data 分片区,可被主参数更新覆盖 |
| 通信与计算 overlap | 难以 overlap(需等待全部梯度才能发起) | 支持 bucket 级别 overlap(overlap_grad_reduce)支持 AG 与前向 overlap( overlap_param_gather) |
| 参数更新后额外通信 | 无需(all-reduce 已同步梯度,参数各 rank 原地更新后一致) | 需 all-gather 重建完整参数(可 overlap 下一轮前向) |
| 综合评价 | 实现简单,无需分片逻辑,但内存高,无法 overlap | 通信量相同,内存大幅节省(12×P 降至 12×P/DP),支持全面 overlap,是大模型标配 |
8.3 何时选择 DistributedOptimizer?
DP ≥ 4 时内存节省已相当可观。
fp32 主参数 + Adam 状态的节省比例为 (1 − 1/DP):
DP=4 节省 75%,DP=8 节省 87.5%,DP=16 节省 93.75%。
对于 70B 以上规模模型,这直接决定了能否在有限显存内完成训练。
模型 fp32 状态超过单卡显存容量时必须使用。 以 7B 参数模型为例(bf16 约 14 GB),fp32 主参数 + Adam 状态约 84 GB, 单张 A100 80G 完全无法容纳,必须依赖 DistributedOptimizer 分摊到多张卡上。
通信带宽充足时,建议同时开启 overlap 选项。
在 NVLink 互联(节点内)或高速 InfiniBand(跨节点)环境下,
overlap_grad_reduce=True + overlap_param_gather=True
可将通信延迟几乎完全隐藏,训练吞吐接近纯计算上限。
带宽有限(如低速以太网)时,overlap 帮助有限,且引入额外代码复杂度,
建议先关闭 overlap 验证正确性。
无需 overlap 或应关闭 overlap 的场景:
调试阶段(overlap 会使错误更难复现)、使用 Pipeline Parallelism 且开启
align_param_gather 时(AG 时机由调度器控制,无需前向 hook)、
或使用 use_megatron_fsdp=True 时(FSDP 路径有独立的 param sync 逻辑)。
8.4 延伸阅读
-
ZeRO 论文(Rajbhandari et al., 2020):
"ZeRO: Memory Optimizations Toward Training Trillion Parameter Models"(SC '20)。
提出了 ZeRO Stage 1/2/3 三个阶段的优化器、梯度和参数分片方案。
Megatron-LM
DistributedOptimizer对应 ZeRO Stage 1(仅分片优化器状态和梯度)。 - DeepSpeed ZeRO: DeepSpeed 的实现支持 ZeRO Stage 1/2/3,Stage 2 额外分片梯度(类比 RS 后不重建), Stage 3 进一步分片模型参数(等价于 FSDP)。与 Megatron-LM 的 DistributedOptimizer 相比,DeepSpeed 的通用性更强但与 Megatron-LM Pipeline/Tensor Parallelism 的深度集成需要额外适配。
-
PyTorch FSDP(Zhao et al., 2023):
"PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel"。
FSDP 对应 ZeRO Stage 3(参数、梯度、优化器状态全部分片),
Megatron-LM 也提供了
use_megatron_fsdp=True选项接入此路径, 内部使用 DTensor 管理分片。 - Megatron-LM 训练效率论文(Narayanan et al., 2021): "Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM", 详细分析了 TP + PP + DP 三维并行组合下的内存模型与通信瓶颈, 是理解 DistributedOptimizer 设计背景的必读材料。
-
源码入口(按阅读顺序):
megatron/core/optimizer/distrib_optimizer.py(优化器主体:分片建立、step、checkpoint)、megatron/core/distributed/param_and_grad_buffer.py(Bucket 管理:start_param_sync、start_grad_sync)、megatron/core/distributed/distributed_data_parallel.py(Hook 注册:backward post-hook、forward pre-hook、overlap 调度)。
内存三块全景与 ZeRO 边界
把混合精度训练里的内存拆成三块,是理解 ZeRO 各阶段边界的最直接方式。
三块内存是什么
nn.Parameter 是它的一个 view。
DP 内完全冗余——同一 PP/TP 位置的所有 DP rank 持有完全相同的内容。
常驻显存
param.main_grad 是它的 view。
Reduce-Scatter 之前:各 rank 持有自己 mini-batch 的局部梯度,内容不同。
Reduce-Scatter 之后:每 rank 只有 1/DP 是有效的全局聚合梯度,但整块 buffer 仍然占着内存。
常驻(到 zero_grad 为止)
fp32 master param:Adam 实际更新的目标,更新后 cast 回 bf16 写入 ①。不是 ① 的备份,是 ① 的源头。
exp_avg (m):一阶矩exp_avg_sq (v):二阶矩
常驻显存
直觉上容易误以为 fp32 master param 是从 bf16 param_data "备份"出来的。 实际相反:每次 step 结束后,是 fp32 master param(Adam 更新后)→ cast → 写回 bf16 param_data。 bf16 是 fp32 的低精度快照,用于高效计算;fp32 才是真正的参数状态。
ZeRO-1 唯一做的事
对比标准 DP,ZeRO-1 只改了一件事:把 All-Reduce 拆成 Reduce-Scatter + All-Gather, 顺手让 optimizer state 跟着分片存。通信总量不变,内存省了 ③ 的 (1 − 1/DP)。
ZeRO-1 / 2 / 3 各自省掉哪块
| 内存块 | 标准 DP | ZeRO-1 | ZeRO-2 | ZeRO-3 |
|---|---|---|---|---|
| ① bf16 param_data | N × 2B | N × 2B | N × 2B | N/DP × 2B ✓ |
| ② bf16 grad_data | N × 2B | N × 2B | N/DP × 2B ✓ | N/DP × 2B ✓ |
| ③ fp32 optimizer state | N × 12B | N/DP × 12B ✓ | N/DP × 12B ✓ | N/DP × 12B ✓ |
| 合计(DP=8,N=4.375B params) | ~61 GB | ~26 GB | ~17 GB | ~9 GB |
注:上表中 N=4.375B 对应 70B 模型在 PP=4, TP=4 下每卡持有的参数量。 标准 DP 合计 = 4.375B × (2+2+12) = 4.375B × 16B ≈ 70 GB; 但 ① 和 ② 约为 8.75 GB 各一份,③ 约 52.5 GB,合计约 70 GB。 ZeRO-1 省掉 ③ 的 7/8 → 省 45.9 GB,合计约 26 GB。 ZeRO-2 再省 ② 的 7/8 → 再省 7.6 GB,合计约 17 GB。 ZeRO-3 再省 ① 的 7/8 → 再省 7.6 GB,合计约 9 GB。
PP 不是 ZeRO-3
PP 和 ZeRO-3 都能让每张卡持有更少的参数,但切法完全不同:
前向前需要 All-Gather 重建完整 W。
DP 内仍然完全冗余——同一 stage 的所有 DP 副本参数完全相同。
PP 减少了"每卡持有的层数",但对于 同一 PP stage 内 的 DP 副本,参数冗余一点没少。 8-way DP 的某个 stage 上,8 张卡的 bf16 param_data 内容完全相同——这正是 ZeRO-3 要解决的问题,PP 无法替代。
Megatron-LM 不做 ZeRO-3 是一个工程权衡:PP+TP 已经把每卡参数量降到足够小(70B → 4.375B), ZeRO-1 再处理 optimizer state,剩余的 bf16 param_data 冗余(约 8.75 GB)相对可接受, 而 ZeRO-3 要求每次前向都 All-Gather bf16 参数,在已有 PP/TP 通信的情况下带宽代价较高。
fp32 main_grad 的特殊地位
前面步骤图里出现的 fp32 main_grad(即 shard_main_param.grad)
并不属于上面三块常驻内存,它是一个临时变量:
# distrib_optimizer.py:2442
shard_main_param.grad = shard_model_grad.float() # ← 新分配 fp32 tensor(N/DP × 4B)
# ↑
# .float() 不是 view,是真实内存分配
分配 fp32 main_grad
.grad = None 释放
它的大小是 N/DP × 4 bytes(已经 ZeRO 分片后的),生命周期只在 optimizer step 期间。
峰值显存需要把它算进去,但稳态(step 完成后)它不存在。