torch.distributed 是 PyTorch 分布式通信核心库,多用于 DDP(DistributedDataParallel)多卡多机训练,负责进程间张量通信、进程组管理。
注意:这是底层通信库,日常训练大多直接用
DistributedDataParallel,但理解底层distAPI 方便做自定义分布式逻辑。
1. 核心概念
- rank:进程编号,
0,1,2...,rank0 常作为主进程 - world_size:总进程数(一张卡一个进程,N卡=world_size=N)
- local_rank:本机内进程编号,本机多张卡从 0 开始
- backend:通信后端
nccl:GPU 分布式,首选,只支持GPUgloo:CPU/GPU,跨平台,GPU性能差mpi:多用于集群MPI环境
2. 完整标准初始化流程
import torch
import torch.distributed as dist
import os
def setup_distributed():
# 1. 从环境变量读取,torchrun 会自动注入这些环境变量
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
# 2. 初始化进程组
dist.init_process_group(
backend="nccl", # GPU用nccl
init_method="env://" # 从环境变量获取主节点地址
)
# 3. 设置当前进程使用的GPU
torch.cuda.set_device(local_rank)
print(f"rank:{rank}, local_rank:{local_rank}, world_size:{world_size}")
def cleanup():
# 销毁进程组
dist.destroy_process_group()
if __name__ == "__main__":
setup_distributed()
# 业务逻辑
cleanup()启动命令(torchrun,推荐)
# 单机2卡
torchrun --nproc_per_node=2 train.py旧版
python -m torch.distributed.launch已经废弃,优先用torchrun。
3. 常用API(四大通信原语)
① dist.all_reduce 所有进程求和/取max,原地修改张量
所有进程张量做规约,结果每一个进程都拿到。
tensor = torch.tensor([1.0, 2.0]).cuda()
# op: SUM, MAX, MIN, PRODUCT
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
# tensor = 所有rank上tensor求和,每个rank都得到总和② dist.all_gather 收集所有进程张量,返回列表
每个进程输出自己的tensor,全部收集到每个进程。
world_size = dist.get_world_size()
local_tensor = torch.tensor([1.0]).cuda()
gather_list = [torch.zeros_like(local_tensor) for _ in range(world_size)]
dist.all_gather(gather_list, local_tensor)
# gather_list[0] rank0数据,gather_list[1] rank1数据...③ dist.broadcast 从rank0广播张量到所有其他进程
rank0的张量复制给全部rank,常用于加载权重、同步参数。
tensor = torch.zeros(3).cuda()
if dist.get_rank() == 0:
tensor = torch.tensor([10.,20.,30.]).cuda()
dist.broadcast(tensor, src=0)
# 所有rank的tensor都变成[10,20,30]④ dist.reduce 规约,只把结果留在dst指定rank
tensor = torch.tensor([1.,2.]).cuda()
dist.reduce(tensor, dst=0, op=dist.ReduceOp.SUM)
# 只有rank0的tensor是总和,其他rank不变⑤ dist.scatter 分发数据,src把张量分片发给各个rank
⑥ dist.barrier() 进程栅栏同步
所有进程走到这里才继续往下跑,用来做同步等待。
dist.barrier()
print("全部进程执行完毕")辅助工具函数
dist.get_rank() # 获取当前进程rank
dist.get_world_size() # 获取总进程数
dist.is_initialized() # 判断进程组是否初始化
dist.is_available() # 判断分布式是否可用4. 最常见场景示例
场景1:分布式训练中求全局loss(all_reduce)
每个GPU算自己的loss,求全局平均loss:
loss = torch.tensor(2.0).cuda()
dist.all_reduce(loss, op=dist.ReduceOp.SUM)
loss = loss / dist.get_world_size()
# rank0打印日志
if dist.get_rank() == 0:
print(f"global loss: {loss.item()}")场景2:只在rank0打印、保存模型
if dist.get_rank() == 0:
print("只有主进程打印日志、保存checkpoint")
torch.save(model.state_dict(), "ckpt.pth")
dist.barrier() # 等待rank0保存完毕,其他进程再继续5. 与 DistributedDataParallel(DDP) 的关系
DDP 内部就是封装了 torch.distributed:
- 前向传播正常跑;反向传播结束后,DDP 自动调用
all_reduce同步梯度,各个进程梯度做all‑reduce求和平均。 - 用户一般不需要手动写
all_reduce梯度。
from torch.nn.parallel import DistributedDataParallel as DDP
model = model.cuda()
model = DDP(model, device_ids=[local_rank])6. 常见踩坑点
- ✅ 通信张量必须放到GPU上(nccl后端,不能传cpu tensor)
- ✅ 所有进程必须执行相同顺序的dist调用,否则死锁
- ✅
all_reduce是原地操作,会修改传入tensor - ✅ 不要混用
DataParallel(DP)和DDP;DP是单进程多线程,不使用torch.distributed - ✅ 多机训练需要设置环境变量
MASTER_ADDR、MASTER_PORT,torchrun自动处理 - ✅ 脚本结束最好调用
dist.destroy_process_group()
7. 进程组分组(subgroup,进阶)
可以创建子进程组做局部通信,比如只对部分卡做allreduce:
group = dist.new_group(ranks=[0,1])
dist.all_reduce(tensor, group=group)