Nvidia 分布式训练配置

1. 分布式训练

  • world_size:参与本次训练的进程总数(通常每个进程对应一张 GPU)。
    它决定了分布式通信组的大小,所有进程必须使用相同的 world_size 才能正确初始化。

  • rank:每个进程的全局唯一标识,取值范围 0 ~ world_size-1。
    在通信中用于识别“我是谁”,比如数据并行时,rank 0 通常承担日志保存、检查点写入等特殊职责。

  • local_rank:进程在当前节点内的编号(0 ~ 节点GPU数-1),主要用于绑定具体的 GPU 设备(torch.cuda.set_device(local_rank))。


1. mpi4py

  • 进程模型:依赖 MPI 环境(mpi4py),通常由 mpirun/mpiexec 预先创建好所有进程,每个进程执行同一脚本。
  • rank / world_size:通过 MPI.COMM_WORLD.Get_rank()Get_size() 获取,完全由 MPI 运行时决定。
  • 进程组初始化:调用 dist.init_process_group,使用外部传入的 dist_url(如 tcp://...)和 MPI 提供的 rank、world_size。
  • 设备绑定local_rank = comm.Get_rank(),然后执行 torch.cuda.set_device(local_rank % num_devices)
  • 适用场景:传统高性能计算集群,需要与 MPI 作业调度系统(如 Slurm + PMI2)集成,或已有庞大 MPI 工作流。

2.torchrun

  • 进程模型:进程由外部 PyTorch 官方启动器 torchrun(或旧版 torch.distributed.launch)预先创建,该启动器负责设置所有必要的环境变量。代码内不再创建子进程。
  • rank / world_size:函数调用时传入 local_rank=-1num_processes=-1dist.init_process_group 会从环境变量 RANKWORLD_SIZE 中读取。dist_url=None 则使用 MASTER_ADDRMASTER_PORT 构建 env:// 初始化方法。

需要给定nnodes/nproc_per_node,NNODES/LOCAL_WORLD_SIZE 等变量,集群环境会设置好

torchrun --nnodes=${NNODES} --nproc_per_node=${LOCAL_WORLD_SIZE} --node_rank=${NODE_RANK} --master_port=${MASTER_PORT}  --master_addr=${MASTER_ADDR} \${XTOUR_DIR}/tools/train.py --config projects/sparse4d_fusion/fvnet_2_2/configs/Sparse4D_henet_fv_virtualcam_v220_fusion.py --stage float --launcher torch

3.mp.spawn

  • 进程模型:在当前进程内通过 torch.multiprocessing.spawn 动态生成 num_processes 个子进程,每个子进程执行 _main_func,并由 spawn 机制依次传入 local_rank(0 到 nprocs-1)。
  • rank / world_size:子进程内 local_rank 由 spawn 提供,world_size=num_processes 固定,dist_url 必须由用户显式传入(如 tcp://127.0.0.1:23456env://)。
  • 进程组初始化:直接使用参数中的 dist_urlnum_processeslocal_rank 调用 dist.init_process_group
  • 设备绑定local_rank != -1,直接用传入的 local_rank 做 local_rank % num_devices 并设置设备。
  • 生命周期管理:主进程捕获 KeyboardInterrupt 后会强制杀死所有子进程(os.killpg),避免孤儿进程。
  • 适用场景:简单的单机多卡训练,不需要额外安装或配置 torchrun/MPI,适合快速原型和轻量级使用。

torch.multiprocessing.spawn

torch.multiprocessing.spawn 是 PyTorch 提供的一个便捷函数,用于启动多个进程来并行执行同一个任务。它尤其常用于多GPU的分布式训练场景。

torch.multiprocessing.spawn(fn, args=(), nprocs=1, join=True, daemon=False)

各参数的含义是:

  • fn:每个子进程都会执行的函数。该函数必须定义在模块的顶层,以便可以被 pickle 序列化。它的第一个参数会被自动传入进程的索引(rank),后面可以跟其他自定义参数。
  • args:一个元组,包含了要传递给 fn 的额外参数。
  • nprocs:要启动的子进程数量,通常等于可用GPU的数量。
  • join:布尔值。若为 True(默认),主进程会阻塞,等待所有子进程执行完毕;若为 False,主进程会立即返回一个 SpawnContext 对象,用于后续手动控制。
  • daemon:布尔值。若设为 True,则创建守护进程。守护进程在主进程结束后会自动终止。