提升 PyTorch 分布式训练效率是一个系统工程,通常可以从数据层面、模型层面、通信层面、硬件层面和代码层面五个维度进行优化。
数据加载往往是瓶颈,尤其是当 GPU 计算速度很快时。
DataLoader 的多进程加载:num_workers > 0。通常建议设置为 CPU 核心数的一半或相等,但不要超过 GPU 数量太多。pin_memory=True。这会将数据张量固定在页锁定内存中,加速 CPU 到 GPU 的数据传输。__getitem__ 中进行。分布式训练的核心开销在于多卡或多机之间的梯度同步。
torch.cuda.amp (Automatic Mixed Precision)。DistributedDataParallel (DDP) 默认会尝试将梯度同步(AllReduce)与反向传播的计算进行重叠。确保你的代码没有阻塞这种操作。DistributedDataParallel (DDP) 而非 DataParallel (DP):print(loss.item()) 是安全的,但 print(loss) 会触发 GPU 到 CPU 的同步,阻塞训练。.item() 和 .cpu() 调用:只在记录日志或保存模型时调用。torch.backends.cudnn.benchmark = True。这会让 cuDNN 自动寻找最适合当前配置的最快卷积算法(前提是输入尺寸固定)。del variable 配合 torch.cuda.empty_cache()(谨慎使用,通常在 epoch 结束时用)。torchrun(推荐,替代旧的 torch.distributed.launch)。torchrun --nproc_per_node=8 train.py。NCCL 环境变量:export NCCL_SOCKET_IFNAME=eth0 # 指定网卡
export NCCL_IB_DISABLE=0 # 如果支持IB,开启
export NCCL_DEBUG=INFO # 调试时开启,查看通信日志以下是快速排查清单,按优先级排序:
DataParallel 改为 DistributedDataParallel (DDP)。torch.cuda.amp。DataLoader (num_workers, pin_memory)。torch.optim.AdamW 配合 AMP,或 Apex FusedAdam)。import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.cuda.amp import autocast, GradScaler
def setup(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
torch.cuda.set_device(rank)
def cleanup():
dist.destroy_process_group()
def train(rank, world_size):
setup(rank, world_size)
model = YourModel().to(rank)
model = DDP(model, device_ids=[rank])
optimizer = torch.optim.AdamW(model.parameters())
scaler = GradScaler() # AMP 缩放器
# 假设 dataset 已定义
sampler = torch.utils.data.distributed.DistributedSampler(
dataset, num_replicas=world_size, rank=rank
)
loader = torch.utils.data.DataLoader(
dataset, batch_size=64, sampler=sampler, num_workers=4, pin_memory=True
)
for epoch in range(10):
sampler.set_epoch(epoch) # 确保每个 epoch shuffle 不同
for data, target in loader:
data, target = data.to(rank), target.to(rank)
optimizer.zero_grad()
# 混合精度上下文
with autocast():
output = model(data)
loss = criterion(output, target)
# 反向传播
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
cleanup()
if __name__ == "__main__":
world_size = torch.cuda.device_count()
mp.spawn(train, args=(world_size,), nprocs=world_size, join=True)