01

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 变成:

  1. Reduce-Scatter:将各 rank 的梯度汇总,并将结果分散到各 rank, 每个 rank 只得到自己"负责"的那段梯度。
  2. 各 rank 独立更新自己持有的参数分片(优化器 step)。
  3. 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,显存空间可以重新分配给更大的批量或更长的序列。

与 ZeRO-2 / ZeRO-3 的区别

DistributedOptimizer 仅实现了 ZeRO-1(仅分散优化器状态), 不做梯度分散(ZeRO-2)或参数分散(ZeRO-3)。Megatron-LM 已通过张量并行(TP)和流水线并行(PP) 在正交维度上切分了参数,因此无需完整的 ZeRO-3,ZeRO-1 已足够大幅降低显存压力。 这也是 DistributedOptimizer 在模型并行训练中的精确定位。

02

类继承结构

三层继承体系

Megatron-LM 的优化器体系由三个核心类构成,形成严格的继承链: MegatronOptimizerMixedPrecisionOptimizerDistributedOptimizer。 每一层各司其职,职责边界清晰。

classDiagram class MegatronOptimizer { <<abstract>> +optimizer: torch.optim.Optimizer +config: OptimizerConfig +init_state_fn: Callable +get_parameters() List +get_main_grads_for_grad_norm() List +get_grad_stats_parallel_group() ProcessGroup +clip_grad_norm(clip_grad) float +count_zeros() float +scale_loss(loss) Tensor +prepare_grads()* bool +step_with_ready_grads()* bool +zero_grad()* +get_loss_scale()* Tensor +reload_model_params()* +state_dict()* +load_state_dict()* +step()* +sharded_state_dict()* } class MixedPrecisionOptimizer { +grad_scaler: MegatronGradScaler +found_inf: Tensor +_dummy_overflow_buf: Tensor +get_loss_scale() Tensor +reload_model_params() +prepare_grads() bool +step_with_ready_grads() bool +step() Tuple #_unscale_main_grads_and_check_for_nan() bool #_copy_model_grads_to_main_grads()* #_collect_main_grad_data_for_unscaling()* #_copy_main_params_to_model_params()* } class DistributedOptimizer { +model_chunks: List +buffers: List[_ParamAndGradBuffer] +data_parallel_group: ProcessGroup +gbuf_ranges: List[Dict] +model_param_gbuf_map: Dict +shard_fp32_from_float16_groups: List +get_grad_stats_parallel_group() ProcessGroup +step_with_ready_grads() bool +_copy_model_grads_to_main_grads() +_collect_main_grad_data_for_unscaling() +_copy_main_params_to_model_params() +zero_grad() +state_dict() Dict +sharded_state_dict() ShardedStateDict #_build_gbuf_range_map()$ #_build_model_gbuf_param_range_map()$ #_build_model_and_main_param_groups()$ } class Float16OptimizerWithFloat16Params { +float16_groups: List +fp32_from_float16_groups: List +fp32_from_fp32_groups: List +step_with_ready_grads() bool } MegatronOptimizer <|-- MixedPrecisionOptimizer MixedPrecisionOptimizer <|-- DistributedOptimizer MixedPrecisionOptimizer <|-- Float16OptimizerWithFloat16Params

第一层:MegatronOptimizer(抽象基类)

定义于 optimizer.py 第 108 行,是纯抽象基类(继承自 ABC)。 它封装了一个底层 torch.optim.Optimizer 实例,并通过 property 将 stateparam_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 行。它是 DistributedOptimizerFloat16OptimizerWithFloat16Params 的共同父类, 两者的差别在于参数分片方式——前者按 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() 流程骨架:

  1. 调用 prepare_grads():将模型梯度拷贝到 fp32 主梯度,并检测 inf/NaN;
  2. 调用 clip_grad_norm()(若开启);
  3. 调用 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_avgexp_avg_sq),切换为其他优化器时需重新实现 sharded_state_dict()

03

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 机制

代码中存在两处对齐填充,各有不同目的:

  1. 参数起始对齐_pad_start_of_param_if_needed):每个参数在 buffer 中的起始偏移必须是 64 的倍数(对应 128-byte 地址对齐,因为 ≥ 16-bit 精度每元素 2 字节)。这是给 DistributedOptimizer 使用的,普通 DDP 不需要。
  2. 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 层次结构图

grad_data(连续 GPU 内存,numel 个元素)
Bucket 0
P
Bucket 1
P
Bucket 2 (tail)
Bucket 0 · offset=0
param_n
·
param_n-1
·
param_n-2
bucket_indices[0] = (0, bucket_end_padded)
Bucket 1 · offset=bucket_end_padded
param_n-3
·
param_n-4
·
param_n-5
·
param_index_map[param] = (start, end, bucket_id)
Bucket 2 (tail)
param_0
剩余参数自动封桶
P = padding(对齐填充,不属于任何参数)
实际参数数据
Bucket(逆序编排,param_n 对应最后一层)
为什么逆序遍历参数?

反向传播遵循 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
04

参数分片——每个 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 类:四种视角的统一抽象

Rangedistrib_optimizer.py 中定义的轻量数据类,仅持有 startendsize 三个字段, 并提供 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,             # 视角④
}

四种视角详解

① gbuf_world
全局 grad_data 坐标系
参数在整个 grad_data buffer(跨所有 bucket)中的绝对偏移区间。 用于从 grad_data tensor 创建参数梯度 view,以及 copy_main_to_model_params 时的索引。
start = param_world_start(在 grad_data 中的全局 offset)
② gbuf_world_in_bucket
当前 bucket 坐标系
将 gbuf_world 的坐标减去 bucket.offset,转换为在当前 bucket grad buffer 内的相对偏移。 AllGather / ReduceScatter 操作直接在单个 bucket 的 tensor 上执行,因此需要此视角。
= Range(gbuf_world.start - bucket.offset, gbuf_world.end - bucket.offset)
③ gbuf_local
本 rank 持有的分片坐标系
参数在本 rank 负责的 1/DP_size 区间内的偏移。ReduceScatter 后每个 rank 拿到的 tensor 从 0 开始, gbuf_local 直接索引这个小 tensor,用于取出本 rank 要更新的梯度片段,传给本地优化器(Adam 等)。
start = max(0, param_world_start - gbuf_world_range.start)
④ param
原始参数 tensor 坐标系
对应参数 tensor 自身(flatten 后)的哪个子区间被本 rank 持有。 由于切分边界不尊重参数边界,一个参数可能被切成两半,分属相邻 rank, sub_param_range 精确描述本 rank 持有的那一段在参数内的位置, 用于 model → main param 的拷贝和 AllGather 后写回参数。
start = max(0, gbuf_world_range.start - param_world_start)

4 个 DP rank 切割 bucket 的示意图

示例:bucket_size=16,DP=4,每 rank 持有 4 个元素
Bucket grad_data(offset=32):
rank 0[0,4) in bucket
rank 1[4,8) in bucket
rank 2[8,12) in bucket
rank 3[12,16) in bucket
全局 grad_data 坐标:[32,36) [36,40) [40,44) [44,48)
各 rank 的 Range 对象内容(以 param P 跨越 rank 0/1 边界为例,P 在 grad_data 全局位于 [34,39)):
字段 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 个元素]
注:param P 共 5 个元素,在 grad_data 全局位于 [34, 39),横跨 rank 0/1 边界(rank 0 持有全局 [32,36))。
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() 方法存在的核心原因——各坐标系之间的转换本质上是一次偏移量的平移。

05

混合精度主参数分片

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 上的特殊标记属性
  • 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_datagrad_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 恢复完整参数。

06

完整训练步骤追踪

DP=4、单个 bucket 含 1024 个 bf16 元素 为例, 追踪从反向传播到下一次前向的完整数据流。 每个 rank 拥有 1024 / 4 = 256 个元素的分片。

DP=4  ·  bucket = 1024 bf16 elements  ·  每 rank 分片 = 256 elements
1
反向传播 → 梯度写入 grad_data buffer
每个参数的 param.main_gradgrad_data buffer 中的一个视图(view),反向传播时梯度直接就地写入连续内存, 无需额外拷贝。DDP 的 hook(param.grad_added_to_main_grad 标志控制) 确保梯度累加到正确位置。
grad_data[0..1023] : bf16 全部 4 ranks 各自本地计算,数据尚未同步
2
Reduce-Scatter 梯度同步
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 元素的已归约梯度
Rank 0
grad[0..1023]
Rank 1
grad[0..1023]
Rank 2
grad[0..1023]
Rank 3
grad[0..1023]
▼ Reduce-Scatter (sum across 4 ranks)
Rank 0
grad[0..255]
sum of 4
Rank 1
grad[256..511]
sum of 4
Rank 2
grad[512..767]
sum of 4
Rank 3
grad[768..1023]
sum of 4
结果原地写回 grad_data buffer 对应分片位置。通信量:每 rank 发送 1024 bf16,接收 256 bf16(已归约)。
3
梯度拷贝到 fp32 主参数(bf16 → fp32)
_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
grad_data[rank×256 .. (rank+1)×256] : bf16 shard_main_param.grad [256] : fp32
4
优化器步骤(纯 fp32 域,无通信)
clip_grad_norm() 计算全局 grad norm(跨 TP/DP rank all-reduce), 随后 Adam.step() 在本 rank 的 256 个 fp32 主参数上执行更新。 各 rank 完全独立执行,不涉及任何通信。
本 rank 的优化器状态(以 Rank 0 为例,256 elements)
param [256] fp32 + exp_avg [256] fp32 + exp_avg_sq [256] fp32 = 256×3×4 bytes = 3 KB
对比非分布式优化器:每 rank 持有 1024×3×fp32 = 12 KB(多 4×)。 DistributedOptimizer 使 Adam 状态随 DP 分片,总内存节省 DP_size 倍
5
fp32 主参数写回 bf16 param buffer(fp32 → bf16)
_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(自动截断精度)
shard_main_param [256] : fp32 param_data[rank×256 .. (rank+1)×256] : bf16
注意坐标系切换:梯度拷贝(阶段 3)使用 param_range(参数内子区间), 参数写回(阶段 5)使用 gbuf_world_in_bucket(全局 bucket 内绝对偏移)。 两套坐标系由 _build_model_gbuf_param_range_map() 统一预计算。
6
All-Gather 参数 → 所有 rank 得到完整参数
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 的视图, 因此自动得到更新,无需额外拷贝。
Rank 0
param[0..255]
already updated
Rank 1
param[256..511]
already updated
Rank 2
param[512..767]
already updated
Rank 3
param[768..1023]
already updated
▼ All-Gather
Rank 0
param[0..1023] bf16
Rank 1
param[0..1023] bf16
Rank 2
param[0..1023] bf16
Rank 3
param[0..1023] bf16
若开启 overlap_param_gather=True,All-Gather 延迟到下一次前向的 pre-hook 中发起,与计算重叠,进一步隐藏通信延迟。
完整调用链(MixedPrecisionOptimizer.step 视角)
step() prepare_grads() _copy_model_grads_to_main_grads() clip_grad_norm() step_with_ready_grads() optimizer.step() _copy_main_params_to_model_params() start_param_sync()
为什么 Reduce-Scatter 之后梯度"自动"在正确位置?

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 中的起点开始。

overlap_grad_reduce 与 overlap_param_gather 两个 overlap 开关

上述 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 负责。

这两个开关不改变数据正确性,只改变通信与计算的调度时序, 是大规模训练的重要吞吐量优化手段。

fp16 场景下的额外步骤:grad scaler 与 inf/nan 检测

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 在是否跳过步骤上达成一致。

07

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 将这一步拆分为:

  1. step 结束时,异步发起第一个 bucket 的 all-gather(而非同步等待完成);
  2. 前向传播中,每个子模块执行前,forward pre-hook 等待该模块所需参数所在 bucket 的 all-gather 完成,同时异步发起下一个 bucket 的 all-gather;
  3. 这样 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(同时开启两个选项)时,单次训练迭代的时间线对比:

无 Overlap(串行)
BW
反向传播
等所有梯度就绪
同步 Reduce-Scatter
RS
Reduce-Scatter
阻塞等待完成
完成通知
OPT
optimizer.step()
更新全量 fp32 分片
同步 All-Gather
AG
All-Gather
阻塞等待完成
完成后才能启动前向
FW
下一次前向
有 Overlap(两个选项均开启)
── 反向传播期间,bucket 完成即触发 RS ──
BW
Bucket N
计算梯度
Bucket 1
计算梯度
Bucket 0
计算梯度
RS
Bucket N RS
异步进行
Bucket 1 RS
异步进行
Bucket 0 RS
异步进行
OPT
optimizer.step()
等所有 RS handle 完成后执行
── step 结束,异步发起第一个 AG,与前向流水线 ──
FW
模块 0 前向
wait AG[0]
模块 1 前向
wait AG[1]
模块 N 前向
wait AG[N]
AG
Bucket 0 AG
异步进行
Bucket 1 AG
异步进行
Bucket N AG
异步进行
BW 反向传播
RS Reduce-Scatter
OPT optimizer.step
AG All-Gather
有 overlap 的阶段同时在时间轴上并列
使用 Overlap 的注意事项

1. handle 必须在使用参数前等待完成。 finish_param_sync() 负责调用 param_gather_handle.wait()。 若跳过此步骤(如在 align_param_gatheroverlap_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 隔离实现异步)。

08

对比与总结

8.1 内存占用对比

设模型参数量为 P,数据并行度为 DPFloat16OptimizerWithFloat16Params(即非分布式优化器)在每个 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 调度)。
09

内存三块全景与 ZeRO 边界

把混合精度训练里的内存拆成三块,是理解 ZeRO 各阶段边界的最直接方式。

三块内存是什么

① bf16 param_data
模型参数本体
前向、反向计算直接用这块内存。 每个 nn.Parameter 是它的一个 view。

DP 内完全冗余——同一 PP/TP 位置的所有 DP rank 持有完全相同的内容。
大小:N × 2 bytes
常驻显存
② bf16 grad_data
参数的梯度
反向传播写入,param.main_grad 是它的 view。

Reduce-Scatter 之前:各 rank 持有自己 mini-batch 的局部梯度,内容不同。
Reduce-Scatter 之后:每 rank 只有 1/DP 是有效的全局聚合梯度,但整块 buffer 仍然占着内存。
大小:N × 2 bytes
常驻(到 zero_grad 为止)
③ fp32 optimizer state
Adam 三件套
fp32 master param:Adam 实际更新的目标,更新后 cast 回 bf16 写入 ①。不是 ① 的备份,是 ① 的源头。

exp_avg (m):一阶矩
exp_avg_sq (v):二阶矩
大小:N × 12 bytes(3 × fp32)
常驻显存
因果方向:fp32 是主,bf16 是导出

直觉上容易误以为 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)。

标准 DP(All-Reduce)
backward → grad_data(局部)
All-Reduce(全量梯度求和广播)
grad_data(全局,每卡相同)
Adam.step() 更新全量 ③
fp32 master + m + v(全量,每卡相同)
cast → 写回
bf16 param_data 更新完毕
ZeRO-1(Reduce-Scatter + All-Gather)
backward → grad_data(局部)
Reduce-Scatter(每卡只得到 1/DP 的聚合梯度)
grad_data 分片(1/DP,各卡不同)
Adam.step() 只更新本卡的 1/DP ③
fp32 master + m + v1/DP,各卡不同)
cast → 写回分片 → All-Gather
bf16 param_data 更新完毕(完整)

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 都能让每张卡持有更少的参数,但切法完全不同:

ZeRO-3:横向切(DP 维度)
同一层权重 W[H, 4H]:
rank 0: W 的 [0, H/4) 列
rank 1: W 的 [H/4, H/2) 列
rank 2: W 的 [H/2, 3H/4) 列
rank 3: W 的 [3H/4, H) 列
每 rank 持有所有层,但每层只有 1/DP。
前向前需要 All-Gather 重建完整 W。
PP:纵向切(模型深度)
Stage 0: Layer 0~7 完整参数
Stage 1: Layer 8~15 完整参数
Stage 2: Layer 16~23 完整参数
Stage 3: Layer 24~31 完整参数
每 rank 只持有部分层,但每层参数是完整的。
DP 内仍然完全冗余——同一 stage 的所有 DP 副本参数完全相同。
PP 不解决 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,是真实内存分配
zero_grad()
... backward ...
prepare_grads()
分配 fp32 main_grad
Adam.step()
zero_grad()
.grad = None 释放

它的大小是 N/DP × 4 bytes(已经 ZeRO 分片后的),生命周期只在 optimizer step 期间。 峰值显存需要把它算进去,但稳态(step 完成后)它不存在。