torch.distributed 是 PyTorch 分布式通信核心库,多用于 DDP(DistributedDataParallel)多卡多机训练,负责进程间张量通信、进程组管理。

注意:这是底层通信库,日常训练大多直接用 DistributedDataParallel,但理解底层 dist API 方便做自定义分布式逻辑。

1. 核心概念

  • rank:进程编号,0,1,2...,rank0 常作为主进程
  • world_size:总进程数(一张卡一个进程,N卡=world_size=N)
  • local_rank:本机内进程编号,本机多张卡从 0 开始
  • backend:通信后端
    • nccl:GPU 分布式,首选,只支持GPU
    • gloo: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. 常见踩坑点

  1. 通信张量必须放到GPU上(nccl后端,不能传cpu tensor)
  2. ✅ 所有进程必须执行相同顺序的dist调用,否则死锁
  3. all_reduce原地操作,会修改传入tensor
  4. ✅ 不要混用 DataParallel(DP)DDP;DP是单进程多线程,不使用torch.distributed
  5. ✅ 多机训练需要设置环境变量 MASTER_ADDRMASTER_PORT,torchrun自动处理
  6. ✅ 脚本结束最好调用 dist.destroy_process_group()

7. 进程组分组(subgroup,进阶)

可以创建子进程组做局部通信,比如只对部分卡做allreduce:

group = dist.new_group(ranks=[0,1])
dist.all_reduce(tensor, group=group)