一、大模型训练的三大瓶颈

1.1 显存瓶颈

以 GPT-3 175B 为例:

组件 显存占用
模型参数 (FP16) 350 GB
Adam 优化器状态 1400 GB
梯度 (FP16) 350 GB
激活值(取决于 batch) 500+ GB
总计 2600+ GB

单张 Ascend 910B 显存 64GB,需要 40+ 张卡才能放下模型,还没算优化器和激活值。

1.2 带宽瓶颈

参数更新时,优化器需要读写全部参数。对 175B 参数,每次更新需要读写 3.5TB 数据,HBM 带宽 1.6TB/s,光参数更新就要 2 秒以上。

1.3 计算瓶颈

矩阵乘法的算术强度低,大部分时间花在数据搬运上,NPU 利用率不到 30%。


二、ZeRO 分片策略

2.1 ZeRO-1: 优化器状态分片

import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP


class ZeRO1Optimizer:
    """ZeRO-1 优化器状态分片

    原理: 每个 rank 只保存 1/N 的优化器状态(Adam 的 m 和 v),
    参数和梯度仍然全量存储。

    显存节省: 优化器状态从 1400GB 降到 35GB/卡(40 卡)
    通信开销: 每次更新需要 AllGather 一次完整参数
    """

    def __init__(self, model, lr=1e-4):
        self.model = model
        self.rank = dist.get_rank()
        self.world_size = dist.get_world_size()

        # 将参数分成 N 份,每个 rank 只优化自己负责的部分
        self.param_groups = self._partition_params()
        self.optimizers = []
        for group in self.param_groups:
            self.optimizers.append(
                torch.optim.Adam(group, lr=lr)
            )

    def _partition_params(self):
        """将参数按 rank 分片"""
        all_params = list(self.model.parameters())
        chunk_size = len(all_params) // self.world_size
        start = self.rank * chunk_size
        end = start + chunk_size if self.rank < self.world_size - 1 else len(all_params)
        return [all_params[start:end]]

    def step(self):
        """执行一步优化

        流程:
        1. 每个 rank 只更新自己负责的参数
        2. AllGather 广播更新后的参数
        3. 所有 rank 同步参数
        """
        # 1. 反向传播得到梯度
        # (梯度在 DDP 中已经自动 AllReduce)

        # 2. 每个 rank 只更新自己的参数
        for opt in self.optimizers:
            opt.step()

        # 3. AllGather 同步参数
        self._allgather_params()

    def _allgather_params(self):
        """AllGather 同步所有 rank 的参数"""
        for group in self.param_groups:
            for param in group:
                dist.all_gather(tensor_list=list(param.chunk(self.world_size)),
                               tensor=param)

2.2 ZeRO-2: 梯度分片

class ZeRO2Optimizer:
    """ZeRO-2 优化器状态 + 梯度分片

    比 ZeRO-1 进一步: 梯度也不需要全量存储。
    每个 rank 只保留自己负责参数的梯度。

    显存节省: 梯度从 350GB 降到 8.75GB/卡
    通信开销: ReduceScatter 替代 AllReduce
    """

    def __init__(self, model, lr=1e-4):
        self.model = model
        self.rank = dist.get_rank()
        self.world_size = dist.get_world_size()

    def backward_step(self, loss):
        """反向传播 + 梯度分片

        梯度计算完成后,立即 ReduceScatter:
        - 每个 rank 收到自己负责参数的梯度和
        - 不需要存储完整梯度
        """
        loss.backward()

        # ReduceScatter: 每个 rank 得到 1/N 的梯度和
        for param in self.model.parameters():
            if param.grad is not None:
                dist.reduce_scatter(
                    output=param.grad[:len(param.grad) // self.world_size],
                    input_list=list(param.grad.chunk(self.world_size)),
                    op=dist.ReduceOp.SUM
                )

2.3 ZeRO-3: 参数 + 梯度 + 优化器全分片

class ZeRO3Strategy:
    """ZeRO-3 全分片策略

    参数、梯度、优化器状态全部分片。
    每个 rank 只保存 1/N 的所有数据。

    显存: 理论最低(总显存 / N)
    通信: 最频繁(前向/反向都要 AllGather 参数)

    适用场景: 模型太大,任何单卡都放不下
    """

    def __init__(self, model):
        self.model = model
        self.rank = dist.get_rank()
        self.world_size = dist.get_world_size()

    def forward_hook(self, module, input, output):
        """前向传播时 AllGather 参数

        每个 layer 前向计算前,先 AllGather 完整参数。
        计算完后立即释放非本地分片。
        """
        for param in module.parameters():
            if param is not None:
                # AllGather 完整参数
                gathered = [torch.zeros_like(param) for _ in range(self.world_size)]
                dist.all_gather(gathered, param.data)
                param.data = torch.cat(gathered, dim=0)

    def cleanup_after_forward(self, module):
        """前向传播后释放非本地参数分片"""
        for param in module.parameters():
            if param is not None:
                # 只保留自己负责的 1/N
                chunk_size = len(param.data) // self.world_size
                start = self.rank * chunk_size
                end = start + chunk_size
                param.data = param.data[start:end]

三、混合精度训练

3.1 AMP 实现

class MixedPrecisionTrainer:
    """混合精度训练管理器

    策略:
    - 前向/反向: FP16 计算(速度快,显存省)
    - 参数更新: FP32 主权重(精度高,稳定)
    - Loss Scaling: 防止 FP16 梯度下溢

    为什么需要 FP32 主权重?
    - FP16 精度只有 3 位指数,更新量太小时会被舍入
    - 累积多次微小更新后,参数偏离真实值
    - FP32 保留一份高精度副本用于更新
    """

    def __init__(self, model, loss_scale_init=2**16, scale_growth_interval=2000):
        self.model = model
        self.loss_scale = loss_scale_init
        self.scale_growth_interval = scale_growth_interval
        self.steps_since_last_scale = 0

        # FP32 主权重
        self.master_weights = {}
        for name, param in model.named_parameters():
            self.master_weights[name] = param.data.float().clone()

    def forward(self, input_data, labels):
        """混合精度前向传播"""
        # 转换输入为 FP16
        input_fp16 = input_data.half()

        # FP16 前向
        with torch.cuda.amp.autocast():
            output = self.model(input_fp16)
            loss = torch.nn.functional.cross_entropy(output, labels)

        # Loss Scaling
        scaled_loss = loss * self.loss_scale
        return scaled_loss, loss

    def backward(self, scaled_loss):
        """混合精度反向传播"""
        scaled_loss.backward()

        # 检查梯度是否 Inf/NaN
        has_overflow = False
        for param in self.model.parameters():
            if param.grad is not None:
                if torch.isinf(param.grad).any() or torch.isnan(param.grad).any():
                    has_overflow = True
                    break

        if has_overflow:
            # 梯度溢出,跳过更新,降低 loss scale
            self.loss_scale /= 2
            self.model.zero_grad()
            return False

        # 调整 loss scale
        self.steps_since_last_scale += 1
        if self.steps_since_last_scale >= self.scale_growth_interval:
            self.loss_scale *= 2
            self.steps_since_last_scale = 0

        return True

    def optimizer_step(self, optimizer):
        """FP32 参数更新"""
        # 梯度反 scaling
        for param in self.model.parameters():
            if param.grad is not None:
                param.grad.data /= self.loss_scale

        # 更新 FP32 主权重
        for name, param in self.model.named_parameters():
            if param.grad is not None:
                self.master_weights[name] -= optimizer.defaults['lr'] * param.grad.float()

        # 将更新后的权重写回 FP16 模型
        for name, param in self.model.named_parameters():
            param.data = self.master_weights[name].half()

四、流水线并行

4.1 GPipe 流水线

class GPipePipeline:
    """GPipe 流水线并行

    将模型按层分成 N 个阶段,每个阶段放在不同的 NPU 上。

    micro-batch 策略:
    - 将一个大 batch 拆成多个 micro-batch
    - 不同 micro-batch 在不同阶段上并行执行
    - 减少流水线气泡

    例子 (4 阶段, 8 个 micro-batch):
    阶段 1: [MB1] [MB2] [MB3] [MB4] [MB5] [MB6] [MB7] [MB8]
    阶段 2:       [MB1] [MB2] [MB3] [MB4] [MB5] [MB6] [MB7] [MB8]
    阶段 3:             [MB1] [MB2] [MB3] [MB4] [MB5] [MB6] [MB7] [MB8]
    阶段 4:                   [MB1] [MB2] [MB3] [MB4] [MB5] [MB6] [MB7] [MB8]
    """

    def __init__(self, model_stages, num_micro_batches=8):
        self.stages = model_stages  # 每个 rank 一个阶段
        self.num_micro_batches = num_micro_batches
        self.rank = dist.get_rank()
        self.world_size = dist.get_world_size()

    def forward_backward(self, input_batch, labels):
        """流水线前向+反向

        1. 前向: 逐 micro-batch 从 stage 0 传到 stage N-1
        2. 反向: 逐 micro-batch 从 stage N-1 传回 stage 0
        3. 参数更新: 每个 stage 只更新自己负责的参数
        """
        # 前向传播
        activations = []
        output = input_batch.chunk(self.num_micro_batches)

        for i, micro_batch in enumerate(output):
            # 发送/接收 activation
            if self.rank > 0:
                micro_batch = self._recv_activation(i)

            # 本地前向
            stage_output = self.stages[self.rank](micro_batch)
            activations.append(stage_output)

            # 发送 activation 到下一阶段
            if self.rank < self.world_size - 1:
                self._send_activation(stage_output, i)

        # 反向传播 (逆序)
        losses = []
        for i in range(self.num_micro_batches - 1, -1, -1):
            # 接收上游梯度
            if self.rank < self.world_size - 1:
                grad = self._recv_gradient(i)

            # 本地反向
            stage_loss = self.stages[self.rank].backward(grad)
            losses.append(stage_loss)

            # 发送梯度到上游
            if self.rank > 0:
                self._send_gradient(stage_loss, i)

        return sum(losses) / len(losses)

    def _send_activation(self, tensor, micro_batch_id):
        """发送 activation 到下一阶段"""
        dist.send(tensor, dst=self.rank + 1)

    def _recv_activation(self, micro_batch_id):
        """从上一阶段接收 activation"""
        tensor = torch.zeros_like(torch.randn(1))
        dist.recv(tensor, src=self.rank - 1)
        return tensor

    def _send_gradient(self, tensor, micro_batch_id):
        dist.send(tensor, dst=self.rank - 1)

    def _recv_gradient(self, micro_batch_id):
        tensor = torch.zeros_like(torch.randn(1))
        dist.recv(tensor, src=self.rank + 1)
        return tensor

五、组合优化策略

def setup_training(model_name="gpt3-175b"):
    """根据模型大小选择最优训练配置"""

    configs = {
        "gpt3-1.5b": {
            "zero_stage": 1,
            "pipeline_parallel": 1,
            "micro_batch": 16,
            "precision": "fp16",
        },
        "gpt3-13b": {
            "zero_stage": 2,
            "pipeline_parallel": 2,
            "micro_batch": 8,
            "precision": "fp16",
        },
        "gpt3-175b": {
            "zero_stage": 3,
            "pipeline_parallel": 8,
            "micro_batch": 4,
            "precision": "fp16",
        },
        "llama-70b": {
            "zero_stage": 2,
            "pipeline_parallel": 4,
            "micro_batch": 4,
            "precision": "bf16",
        },
    }

    config = configs.get(model_name, configs["gpt3-175b"])
    print(f"模型 {model_name} 训练配置:")
    for k, v in config.items():
        print(f"  {k}: {v}")
    return config

六、常见问题

问题 原因 解决方案
显存还是不够 ZeRO 分片数太少 增加 DP 并行度,启用 ZeRO-3
训练速度太慢 流水线气泡太大 增加 micro-batch 数量
训练不收敛 Loss scale 不合适 自动调整 loss scale
梯度 NaN 学习率太高 降低学习率,warmup 更长

相关仓库

  • CANN - 昇腾计算架构 https://gitee.com/ascend/cann
  • DeepSpeed - ZeRO 优化库 https://github.com/microsoft/DeepSpeed
  • Megatron-LM - 大模型训练框架 https://github.com/NVIDIA/Megatron-LM
  • MindSpore - 华为训练框架 https://github.com/mindspore-ai/mindspore
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐