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=-1和num_processes=-1,dist.init_process_group会从环境变量RANK、WORLD_SIZE中读取。dist_url=None则使用MASTER_ADDR和MASTER_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:23456或env://)。 - 进程组初始化:直接使用参数中的
dist_url、num_processes、local_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,则创建守护进程。守护进程在主进程结束后会自动终止。
