CANN 大模型训练优化:从显存受限到千亿参数的实战路径
·
一、大模型训练的三大瓶颈
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
更多推荐




所有评论(0)