You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何理解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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 15:01:05