大模型训练分布式优化ZeRO-3与DeepSpeed Zero阶段显存分配策略实战配置

大模型训练过程中,显存瓶颈是最常遇到的工程问题。一个百亿参数模型用FP16精度存储就需要近200GB显存,单张A100 80GB根本无法装下。ZeRO(Zero Redundancy Optimizer)通过将优化器状态、梯度和参数分片到不同设备上,让多卡训练成为可能。DeepSpeed框架实现了ZeRO的三个阶段,合理配置可以将万亿参数模型训练压缩到可接受的硬件规模内。

ZeRO三阶段显存优化原理与参数分片策略

ZeRO的核心思想是消除数据并行中的显存冗余。传统数据并行中,每张卡都保存完整的模型参数、梯度和优化器状态,造成大量重复占用。ZeRO将这三部分逐步分片:

ZeRO-1(优化器状态分片):仅将Adam优化器的状态(一阶矩、二阶矩)按卡切分,每张卡只保留1/N的优化器状态。对于AdamW,优化器状态约占参数量的2倍FP16空间,分片后单卡显存占用显著降低。

ZeRO-2(梯度分片):在ZeRO-1基础上,将梯度也按卡切分。每张卡只负责计算和更新自己分片对应的梯度,减少梯度通信量。反向传播过程中通过reduce-scatter操作完成梯度聚合与分片同步。

ZeRO-3(参数分片):将模型参数本身也按卡切分。前向传播时按需通过all-gather动态收集对应层的参数,用完立即释放。这带来额外的通信开销,但显存占用降至最低,使超大模型训练可行。

DeepSpeed ZeRO-3配置文件与参数详解

DeepSpeed通过JSON配置文件控制ZeRO行为。以下是一个典型的ZeRO-3训练配置:

{
  "train_micro_batch_size_per_gpu": 4,
  "gradient_accumulation_steps": 8,
  "zero_optimization": {
    "stage": 3,
    "overlap_comm": true,
    "contiguous_gradients": true,
    "sub_group_size": 1e9,
    "reduce_bucket_size": 5e8,
    "stage3_prefetch_bucket_size": 5e8,
    "stage3_param_persistence_threshold": 1e5,
    "stage3_max_live_parameters": 1e9,
    "stage3_max_reuse_distance": 1e9,
    "stage3_gather_16bit_weights_on_model_save": true
  },
  "fp16": {
    "enabled": true,
    "loss_scale": 0,
    "loss_scale_window": 1000,
    "initial_loss_scale": 65536,
    "hysteresis": 2,
    "min_loss_scale": 1
  },
  "optimizer": {
    "type": "AdamW",
    "params": {
      "lr": 2e-5,
      "betas": [0.9, 0.95],
      "eps": 1e-8,
      "weight_decay": 0.01
    }
  },
  "activation_checkpointing": {
    "partition_activations": true,
    "cpu_checkpointing": false,
    "contiguous_memory_optimization": true,
    "number_checkpoints": 4,
    "synchronize_checkpoint_boundary": false
  },
  "gradient_clipping": 1.0
}

关键参数说明:

overlap_comm:设为true时,参数通信与计算重叠执行,隐藏通信延迟。生产环境建议开启,但会增加约10%的显存开销用于缓冲区。

reduce_bucket_size:控制梯度归约的桶大小。值越大通信次数越少但显存占用越高,5e8是一个在A100上经过验证的合理值。

stage3_prefetch_bucket_size:前向传播时预取参数的桶大小,影响计算与通信的重叠效率。

stage3_param_persistence_threshold:小于此阈值的参数不分片,常驻每张卡。小参数分片反而增加通信开销,设为1e5可避免这个问题。

ZeRO-3与PyTorch Lightning集成训练代码示例

以下代码展示如何在PyTorch中使用DeepSpeed ZeRO-3训练一个GPT类模型:

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, get_linear_schedule_with_warmup
from torch.utils.data import DataLoader
import deepspeed

def create_model_and_optimizer(model_path, config_path):
    model = AutoModelForCausalLM.from_pretrained(
        model_path,
        torch_dtype=torch.float16
    )
    
    tokenizer = AutoTokenizer.from_pretrained(model_path)
    
    # DeepSpeed初始化
    model_engine, optimizer, _, _ = deepspeed.initialize(
        model=model,
        model_parameters=model.parameters(),
        config=config_path
    )
    
    return model_engine, optimizer, tokenizer

def train_epoch(model_engine, dataloader, scheduler, device_id):
    model_engine.train()
    total_loss = 0
    
    for step, batch in enumerate(dataloader):
        batch = {k: v.to(model_engine.local_rank) for k, v in batch.items()}
        
        outputs = model_engine(**batch)
        loss = outputs.loss
        
        # 反向传播由DeepSpeed引擎处理
        model_engine.backward(loss)
        
        # 梯度累积逻辑
        if step % model_engine.gradient_accumulation_steps() == 0:
            model_engine.step()
            scheduler.step()
            model_engine.zero_grad()
        
        total_loss += loss.item()
        
        if step % 100 == 0 and model_engine.local_rank == 0:
            print(f"Step {step}, Loss: {loss.item():.4f}, "
                  f"LR: {scheduler.get_last_lr()[0]:.2e}")
    
    return total_loss / len(dataloader)

# 启动训练
if __name__ == "__main__":
    model_engine, optimizer, tokenizer = create_model_and_optimizer(
        "/data/models/llama-13b",
        "ds_config_zero3.json"
    )
    
    scheduler = get_linear_schedule_with_warmup(
        optimizer,
        num_warmup_steps=500,
        num_training_steps=10000
    )
    
    # 多卡启动
    # torchrun --nproc_per_node=8 train_zero3.py
    for epoch in range(3):
        avg_loss = train_epoch(
            model_engine, dataloader, scheduler, 
            model_engine.local_rank
        )
        if model_engine.local_rank == 0:
            print(f"Epoch {epoch}: avg_loss={avg_loss:.4f}")

ZeRO-3通信开销与计算效率平衡调优

ZeRO-3的主要代价是额外的通信开销。每次前向传播需要all-gather收集参数,反向传播需要额外的reduce-scatter。通信量与模型参数量成正比,在万兆网络下,通信可能成为瓶颈。

实测数据参考(8xA100 80GB,NVLink互联):

13B模型训练,ZeRO-2吞吐量约420 samples/s,ZeRO-3约360 samples/s,性能下降约14%。但ZeRO-3单卡显存占用从约104GB降至约18GB,使得训练成为可能。

70B模型训练,ZeRO-2无法在8卡上运行(显存溢出),ZeRO-3在8卡上吞吐量约48 samples/s。这是可接受的训练速度。

调优建议:

1. 优先使用ZeRO-2,仅在显存不足时切换到ZeRO-3。ZeRO-2的通信开销更小,计算效率更高。

2. 开启overlap_commcontiguous_gradients减少通信等待时间。

3. 调整reduce_bucket_sizeprefetch_bucket_size找到通信次数与显存占用的平衡点。从5e8开始,观察GPU利用率,逐步调整到GPU利用率最高的值。

4. 配合激活检查点(Activation Checkpointing)进一步降低显存。开启后可减少约40%-60%的激活显存占用,代价是约20%-30%的计算时间增加。

5. 如果NVLink带宽不足,考虑使用ZeRO-2而非ZeRO-3。PCIe互联下ZeRO-3的通信开销可能超过50%。

ZeRO-3模型保存与检查点恢复

ZeRO-3下参数分片存储,保存和加载检查点需要特殊处理:

# 保存检查点(每张卡保存自己的分片)
def save_checkpoint(model_engine, output_dir, epoch):
    if model_engine.local_rank == 0:
        os.makedirs(output_dir, exist_ok=True)
    
    # DeepSpeed会自动处理分片保存
    model_engine.save_checkpoint(
        os.path.join(output_dir, f"checkpoint-{epoch}"),
        tag=f"epoch_{epoch}"
    )

# 保存完整16位权重用于推理部署
def save_full_model(model_engine, tokenizer, output_dir):
    if model_engine.local_rank == 0:
        # stage3_gather_16bit_weights_on_model_save需设为true
        # DeepSpeed会自动收集所有分片并保存完整模型
        model_engine.save_16bit_model(
            output_dir,
            save_filename="pytorch_model.bin"
        )
        tokenizer.save_pretrained(output_dir)

# 恢复训练
def resume_training(model_engine, checkpoint_dir):
    _, client_state = model_engine.load_checkpoint(
        checkpoint_dir,
        tag="latest"
    )
    # client_state包含epoch、step等自定义状态
    return client_state.get("epoch", 0), client_state.get("step", 0)

注意:保存完整16位模型时,stage3_gather_16bit_weights_on_model_save必须设为true,否则保存的只是当前卡的分片,无法直接用于推理。如果显存不足以一次性收集所有参数,可以在推理阶段用DeepSpeed的init_inference配合ZeRO-3进行分片推理。

原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/da-mo-xing-xun-lian-fen-bu-shi-you-hua-zero3-yu/

(0)
小编小编
上一篇 19小时前
下一篇 19小时前

相关推荐