如何使用PyTorch在多台虚拟机上实现深度学习模型分布式训练
多节点PyTorch分布式训练方案(适配AWS 3~4台实例场景)
梯度汇总核心说明
- 梯度汇总、跨节点通信逻辑不需要手动实现,现有成熟框架已经做了全封装,直接调用即可
- 关于梯度求和还是取平均:当前工业界通用的分布式数据并行策略默认对所有节点的梯度取平均,逻辑和你之前用的单机DataParallel完全对齐,相当于把所有节点的样本合并为一个全局大批次计算梯度,不需要手动调整参数。
PyTorch原生方案:DistributedDataParallel (DDP)
你之前使用的单机DataParallel是单进程多线程实现,仅支持单节点多卡场景,本身不支持跨节点通信。PyTorch原生提供的torch.nn.parallel.DistributedDataParallel(简称DDP)是目前最稳定的跨节点训练工具,采用多进程实现,每个GPU对应独立进程,通信效率远高于DataParallel,你的场景完全可以直接使用。
DDP代码改造与运行步骤
- 提前配置AWS实例环境:所有实例放在同一VPC、同一可用区下,安全组开放所有实例之间的TCP通信端口(默认用23456即可),保证所有实例的PyTorch版本、Python版本完全一致。
- 原有单机训练代码仅需修改几处:
import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler # 初始化分布式进程组,GPU训练默认用nccl通信后端 dist.init_process_group(backend="nccl") local_rank = dist.get_rank() % torch.cuda.device_count() torch.cuda.set_device(local_rank) # 模型套入DDP wrapper model = YourCustomModel().cuda(local_rank) model = DDP(model, device_ids=[local_rank]) # 数据集用DistributedSampler自动切分,避免多节点加载重复数据 train_dataset = YourTrainDataset() train_sampler = DistributedSampler(train_dataset) train_loader = DataLoader(train_dataset, batch_size=单卡批次大小, sampler=train_sampler) # 训练循环 for epoch in range(total_epochs): train_sampler.set_epoch(epoch) # 保证每个epoch数据shuffle逻辑正确 for batch in train_loader: # 前向传播、损失计算、反向传播逻辑和原有代码完全一致 loss = model(batch) loss.backward() optimizer.step() optimizer.zero_grad()
- 启动训练:所有节点执行同一条启动命令,仅需修改
--node_rank参数,主节点填0,其余节点依次填1、2、3即可:
torchrun --nproc_per_node=每台实例的GPU数量 --nnodes=实例总数 --node_rank=当前节点序号 --master_addr=主节点私有IP --master_port=23456 train.py
注意
master_addr填主节点的AWS内网IP,所有节点的其余启动参数必须完全一致。
PyTorch Lightning简化方案
如果不想手动处理分布式初始化逻辑,PyTorch Lightning已经把所有节点通信、梯度同步、数据切分逻辑做了封装,不需要调用额外的通信模块,仅需修改Trainer初始化参数即可:
from pytorch_lightning import Trainer # 仅需修改这一处配置,其余训练逻辑和单机代码完全一致 trainer = Trainer( accelerator="gpu", devices="auto", num_nodes=4, # 你的实例总数量 strategy="ddp" ) trainer.fit(model)
启动命令和上述原生DDP的torchrun命令完全相同,不需要做额外改造。
AWS场景优化建议
- 优先选择同可用区的g5/p3/p4系列GPU实例,节点内网延迟更低,训练效率更高
- 用EFS或者S3作为共享存储挂载到所有实例,无需每台实例单独复制训练数据集
- 若出现通信报错,优先检查安全组规则是否放开了所有节点之间的指定端口,以及NCCL依赖是否安装完整
内容的提问来源于stack exchange,提问作者Adam Appletree
相关产品推荐
相关产品推荐

