导读:本期聚焦于天穹小白创作的《Megatron-LM与DeepSpeed的分布式训练源码里,模型和优化器到底是怎么切分的?》,敬请观看详情。大模型训练遇到显存瓶颈时,Megatron-LM选择把单层权重按列或按行切到不同GPU,用通信换显存;DeepSpeed则把优化器状态、梯度和参数本身分片存储,用显存换计算。本文直接从源码调用链切入,解析ColumnParallelLinear与RowParallelLinear的前向计算和反向梯度同步过程,说明张量并行如何减少单卡激活值显存占用。随后拆解DeepSpeed ZeRO Stage1与Stage2的分区优化器实现,包括梯度reduce-scatter、参数all-gather以及Adam状态更新。最后分析两者在deepspeed.initialize和Megatron-DeepSpeed适配层中的配合方式,梳理混合并行的初始化顺序。读完能对分布式训练框架的内部机制有具体认识,不再把它们当作黑盒。

拆解Megatron-LM和DeepSpeed的源码,本质上是在回答一个问题:当单卡显存装不下模型或者优化器状态时,框架如何把计算和存储拆到多卡上。Megatron-LM的思路是模型并行,重点在单层内的矩阵乘法切分;DeepSpeed的思路是数据并行下的状态分片,重点在优化器和梯度。两者并不冲突,反而可以在同一个训练任务里叠加使用,这正是很多大模型训练框架的底层逻辑。

Megatron-LM与DeepSpeed的分布式训练源码里,模型和优化器到底是怎么切分的?

一、Megatron-LM张量并行的源码核心:ColumnParallelLinear与RowParallelLinear

张量并行要解决的核心问题是:单个Transformer层中的线性层权重太大,单卡放不下或者计算太慢。Megatron-LM把线性层权重按输出维度或输入维度切分到多个GPU,每张卡只保存权重的一部分,但仍然能通过通信还原完整的计算结果。源码中这部分实现集中在megatron/core/tensor_parallel/layers.py,两个核心类分别是ColumnParallelLinear和RowParallelLinear。

ColumnParallelLinear的构造函数会把out_features除以world_size,所以每张卡只持有输出维度的一部分。在forward中,它先对完整输入做本地线性变换,得到部分输出,然后调用gather_from_tensor_model_parallel_region执行all-gather,把所有卡上的部分输出拼接成完整输出。反向传播时,这个all-gather的梯度会自动变成reduce-scatter,把完整梯度切回各卡对应维度。

下面的简化代码展示了列并行和行并行的前向逻辑,省略了真实实现中的通信调用,但保留了切分维度的关键差异。

import torch
import torch.nn as nn
import torch.nn.functional as F

class ColumnParallelLinear(nn.Module):
    def __init__(self, in_features, out_features, world_size, rank):
        super().__init__()
        self.rank = rank
        self.world_size = world_size
        # 输出维度按world_size切分,每张卡只保留 out_features/world_size 个神经元
        self.weight = nn.Parameter(torch.empty(out_features // world_size, in_features))
        self.bias = nn.Parameter(torch.empty(out_features // world_size))

    def forward(self, x):
        # 输入x在所有GPU上保持一致,这里只计算本地持有的输出分片
        local_output = F.linear(x, self.weight, self.bias)
        # 真实Megatron实现中需要在此处调用 all_gather 把各卡输出拼成完整输出
        # 简化示例略去通信调用
        return local_output

class RowParallelLinear(nn.Module):
    def __init__(self, in_features, out_features, world_size, rank):
        super().__init__()
        self.rank = rank
        self.world_size = world_size
        # 输入维度按world_size切分,每个GPU只保留输入的一部分特征
        self.weight = nn.Parameter(torch.empty(out_features, in_features // world_size))
        self.bias = nn.Parameter(torch.empty(out_features))

    def forward(self, x):
        # x已经在各GPU上按列切分,本地计算得到部分和
        partial_sum = F.linear(x, self.weight, None)
        # 对所有GPU的partial_sum做all-reduce,再加bias得到完整输出
        # 真实实现使用 torch.distributed.all_reduce
        # 简化示例略去通信调用
        output = partial_sum + self.bias
        return output

RowParallelLinear的切分方向正好相反,它把输入维度切成多份,每张卡只保存in_features/world_size列权重。前向时输入x已经按列切分,各卡算出部分和,再通过all-reduce把所有卡的部分和加起来,最后加上完整的bias。这样设计的巧妙之处在于,列并行层后面接行并行层时,列并行层的all-gather输出正好可以作为行并行层的分片输入,减少一次不必要的全局同步。

张量并行的优点是计算效率高,每张卡只算一部分矩阵乘法,计算量线性下降。但缺点是通信非常频繁,一个前向传播里就要发生多次all-gather和all-reduce。因此张量并行通常只用于单机内部的高速NVLink互联,跨节点带宽不足时会成为瓶颈。

二、DeepSpeed ZeRO的优化器状态分片:从初始化到分区更新

DeepSpeed解决的是另一个维度的显存问题。在纯数据并行中,每张卡都保存完整模型副本和完整优化器状态,显存冗余非常严重。以混合精度训练7.5B参数模型为例,FP16参数约15GB,梯度约15GB,Adam状态包括FP32参数、一阶动量、二阶动量各30GB,总计超过120GB。单卡32GB显存完全无法容纳。ZeRO的核心思想是把这些状态在数据并行的多张卡之间分片存储,每张卡只保留自己负责的那一部分。

ZeRO分为三个阶段:Stage1只分片优化器状态,Stage2继续分片梯度,Stage3连参数也分片。源码入口在deepspeed.initialize,它会读取配置中的zero_optimization.stage字段,然后创建对应的分区优化器。下面是最常见的Stage2初始化写法。

import deepspeed

# 配置文件:开启ZeRO Stage2
ds_config = {
    "train_batch_size": 32,
    "gradient_accumulation_steps": 1,
    "zero_optimization": {
        "stage": 2,
        "contiguous_gradients": True,
        "overlap_comm": True
    },
    "fp16": {
        "enabled": True
    }
}

model_engine, optimizer, _, _ = deepspeed.initialize(
    model=model,
    optimizer=optimizer,
    config=ds_config,
    model_parameters=model.parameters()
)

# 之后的标准训练循环
for step, batch in enumerate(dataloader):
    loss = model_engine(batch)
    model_engine.backward(loss)
    model_engine.step()

在DeepSpeedZeroOptimizer_Stage2的step函数中,更新流程和普通优化器完全不同。普通数据并行中,每张卡先对梯度做all-reduce拿到完整平均梯度,然后各自更新完整参数。Stage2则先对梯度做reduce-scatter,让每张卡只保留自己负责那部分参数的完整梯度,然后只更新这些参数对应的FP32主权重和Adam状态,最后再通过all-gather把更新后的参数分片收集成完整参数。这样可以省去保存完整梯度和完整Adam状态的开销。

下面的伪代码梳理了Stage2单次step的关键步骤。

# DeepSpeed ZeRO Stage2 优化器step的简化流程
def zero_stage2_step(self):
    # 1. 对每个参数的梯度做reduce-scatter
    # 原本完整梯度分布在所有GPU上,reduce-scatter后每个GPU只保留一部分参数的完整梯度
    self.reduce_scatter_gradients()

    # 2. 每个GPU只更新自己负责的那部分参数
    # 需要同时更新对应的fp32 master weights、fp16 weights和Adam状态
    for param_group in self.optimizer.param_groups:
        for param in param_group['params']:
            if param.grad is None:
                continue
            # 获取该参数对应的fp32 master weight
            fp32_param = self.fp32_partitioned_groups[param_group['name']][param]
            # 获取梯度分片
            grad = param.grad
            # 调用Adam更新fp32参数
            self.adam_update(fp32_param, grad)
            # 将fp32参数转回fp16并写回param
            param.data.copy_(fp32_param.float())

    # 3. 所有GPU对更新后的参数分片做all-gather,恢复完整参数
    self.all_gather_parameters()

Stage2的显存收益在8卡配置下非常明显。参数仍每卡保存完整15GB,梯度分片后每卡约1.875GB,Adam状态分片后每卡约11.25GB,总显存需求降到约28.125GB,刚好可以放进32GB显卡。Stage3进一步把参数也分片,每卡参数只需要1.875GB,但前向和反向时需要通过all-gather临时恢复完整参数,通信量更大。

这里需要区分一个容易混淆的点:ZeRO的reduce-scatter和all-gather并不是简单的all-reduce替换。在Stage2中,reduce-scatter发生在backward结束后,把梯度从完整状态切成分片状态;all-gather发生在step结束后,把参数从分片状态恢复成完整状态。这两次通信可以用overlap_comm与反向计算重叠,也可以在CPU上执行,以减少对训练吞吐的影响。

三、Megatron-LM与DeepSpeed的协同:混合并行初始化调用链

实际训练大模型时,很少单独使用某一种并行策略。常见的做法是先用Megatron-LM做张量并行和流水线并行,把模型切到单机内的多张卡上,再用DeepSpeed对切分后的模型做ZeRO数据并行,跨机扩展。这样每一层内部通过张量并行减少单卡计算量和激活值占用,层之间通过流水线并行减少存储压力,数据并行维度上再通过ZeRO压缩优化器状态。

两者集成的代码入口在Megatron-DeepSpeed仓库中。启动训练时,需要同时设置tensor-model-parallel-size、pipeline-model-parallel-size和zero-stage等参数。模型构建阶段先按照张量并行和流水线并行切分,然后调用deepspeed.initialize,后者会遍历模型所有参数,根据ZeRO配置创建分区优化器。初始化顺序不能颠倒,因为ZeRO需要知道模型切分后的参数布局,才能正确映射每个参数分片归属。

下面是一段简化的配置文件,展示混合并行参数如何同时传入。

# 混合并行配置示例:张量并行2卡,流水线并行2卡,ZeRO Stage2
parallel_config = {
    "tensor_model_parallel_size": 2,
    "pipeline_model_parallel_size": 2,
    "zero_stage": 2,
    "train_batch_size": 64,
    "micro_batch_size": 4,
    "gradient_accumulation_steps": 8,
    "fp16": {
        "enabled": True,
        "loss_scale": 0,
        "initial_scale_power": 16
    }
}

# 初始化顺序:先初始化分布式进程组,再构建模型并切分,最后调用deepspeed.initialize
torch.distributed.init_process_group(backend='nccl')
model = build_megatron_model(parallel_config)
model_engine, optimizer, _, _ = deepspeed.initialize(
    model=model,
    optimizer=optimizer,
    config=parallel_config,
    model_parameters=model.parameters()
)

从源码调用链来看,deepspeed.initialize内部会执行DeepSpeedEngine.__init__,随后进入_configure_zero_optimizer,根据zero_optimization.stage选择DeepSpeedZeroOptimizer_Stage2或Stage3类。这些类在构造时会读取当前数据并行rank和world size,把优化器状态按参数维度切分。如果模型已经经过张量并行切分,那么每个张量并行组内部的参数分片会被视为一个独立参数,ZeRO只在不同数据并行组之间分片。

这种协同方式的一个关键收益是通信量的均衡。张量并行产生的all-reduce通信只发生在单机内部的张量并行组中,ZeRO产生的reduce-scatter和all-gather只发生在数据并行组中。两者通过不同的通信域隔离,互不干扰。调试时建议先把ZeRO关掉,只开张量并行,确认模型前向反向正确;再逐步开启ZeRO Stage1和Stage2,观察显存变化和loss是否稳定。

四、通信原语选择与性能调优经验

理解源码后会意识到,分布式训练框架的底层其实就是在几种通信原语之间做权衡。all-reduce适合行并行的梯度汇总,all-gather适合列并行输出拼接和ZeRO参数恢复,reduce-scatter适合ZeRO梯度分片。它们在带宽消耗上并不相同,比如N-Cube算法下的all-reduce需要约2倍数据量的跨卡传输,而reduce-scatter和all-gather各需要约1倍数据量的传输,但两者相加又回到2倍。因此Stage2单步的总通信量和普通数据并行接近,只是分成了两个阶段。

调优时最有效的几个参数是overlap_comm、contiguous_gradients和allgather_bucket_size。开启通信重叠后,反向计算和梯度reduce-scatter可以同时进行,隐藏掉一部分通信延迟。contiguous_gradients会把多个小梯度拷贝到连续缓冲区,减少通信启动次数。allgather_bucket_size控制参数收集时的分桶大小,过大容易造成显存峰值,过小则通信次数增多。这些参数需要结合具体模型和硬件做几轮实验才能找到平衡点。

另一个容易被忽略的调优点在张量并行的反向通信。列并行层的反向需要梯度reduce-scatter,行并行层的反向需要梯度all-gather。如果张量并行度设置过高,比如单机8卡全部做张量并行,通信量会急剧增加,反而可能比4卡张量并行加2路数据并行更慢。因此,实践中通常让张量并行度不超过单机GPU数量,剩余卡全部用于数据并行和ZeRO分片。这样既能发挥NVLink的低延迟,又能利用ZeRO的显存压缩。

Megatron-LMDeepSpeed分布式训练修改时间:2026-09-22 19:52:20

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/0922/60606.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。