如何理解PyTorch中local_rank参数及if-else分支的设备设置逻辑?
PyTorch 设备设置代码逻辑解释
local_rank参数的具体含义
- 这个参数是PyTorch分布式训练启动工具(比如
torchrun、旧版的torch.distributed.launch)自动注入到每个进程的参数,代表当前进程在本机对应的GPU编号。如果是单卡/CPU非分布式模式运行脚本,通常会手动把这个参数设为-1,用来标记不需要走分布式逻辑。
两种设备设置方式的差异与代码运行逻辑
你贴出的代码是PyTorch训练脚本里非常经典的「兼容单设备运行+多卡分布式运行」的设备初始化逻辑,代码如下:
if args.local_rank == -1: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') else: torch.distributed.init_process_group(backend='nccl') torch.cuda.set_device(args.local_rank) device = torch.device('cuda', args.local_rank)
非分布式分支(args.local_rank == -1)
- 触发场景:普通方式直接启动脚本,不使用分布式启动工具,单设备运行。
- 运行逻辑:自动检测当前环境是否有可用的NVIDIA GPU,有就默认使用第0张CUDA显卡,没有可用GPU就自动切换到CPU运行。
- 这个模式下不需要做多进程通信初始化,所有运算都跑在单个设备上。
分布式训练分支(args.local_rank != -1)
- 触发场景:用分布式启动工具启动脚本,做多卡(可以是单机多卡,也可以是多机多卡)分布式训练。启动工具会给每个进程分配互不重复的
local_rank值,比如单机4卡就会启动4个进程,对应local_rank值为0、1、2、3,每个进程负责一张卡的运算。 - 逐行运行逻辑:
torch.distributed.init_process_group(backend='nccl'):初始化分布式进程通信组,选nccl作为通信后端是因为它针对NVIDIA GPU做了硬件级优化,多卡间数据传输效率最高,是GPU分布式训练的标准选择。torch.cuda.set_device(args.local_rank):将当前进程和指定编号的GPU绑定,避免所有进程默认都往第0张卡塞数据导致显存溢出。torch.device('cuda', args.local_rank):生成对应绑定GPU的设备对象,后续把模型、张量迁移到这个设备时,就会自动落到当前进程负责的那张GPU上,不会和其他进程的设备冲突。
内容的提问来源于stack exchange,提问作者Wei Zhang
相关产品推荐
相关产品推荐

