大模型训练过程中,显存瓶颈是最常遇到的工程问题。一个百亿参数模型用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_comm和contiguous_gradients减少通信等待时间。
3. 调整reduce_bucket_size和prefetch_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/