跳到主要内容

创建 PyTorchDDP 分布式训练

更新时间:2025-08-19 04:30:57

PyTorch 是科研和工业界广泛使用的主流深度学习框架,以动态图机制为核心,兼具灵活性与易用性,并提供丰富的分布式训练、模型部署和生态工具支持。

PyTorch 的 DDP(Distributed Data Parallel)分布式训练模式通过在多个 GPU 或节点间同步梯度实现高效的数据并行,具有易用、通信开销低、扩展性强等特点,是大规模分布式训练的主流方式。常使用 torchrun 命令启动分布式训练。

前提条件​

  • 选择的镜像中需要已安装 torch。

操作步骤​

  1. 请参考创建分布式训练模块填写基本配置。
  2. 计算框架选择“其他”。
  3. 任务配置请设置一个 worker Role(节点组)即可。
  4. 填写节点数量。
  5. 输入启动命令,参考后文使用方式。

使用 torchrun 启动训练​

torchrun 是 PyTorch 推荐的分布式训练启动器,用于简化多 GPU、多节点训练任务的启动与管理;它的优势是配置简单(无需手动指定 node rank 等参数)、支持弹性训练(节点可动态加入或退出)、并且统一了分布式脚本的运行方式。

以“一个 worker role 共2个节点,每个节点 8 个 GPU” 为例,输入启动命令如下:

# 取 worker 节点组中第 1 个节点作为 rendezvous 通信的 master 节点
export MASTER=`echo ${VC_WORKER_HOSTS} | awk -F , '{print $1}'`

torchrun \
--rdzv-id=0 \
--rdzv-backend=c10d \
--rdzv-endpoint=${MASTER}:29500 \
--nnodes=2 \
--nproc_per_node=8 \
YOUR_TRAINING_SCRIPT.py \
<script args>

其中:

  • 环境变量 VC_WORKER_HOSTS 是当前 worker role 所有节点的地址,以逗号分隔。通过以上 shell 语句获取第一个节点的地址,赋值给 MASTER 环境变量。
  • --rdzv-id=0是自定义的当前训练任务的名字。
  • --rdzv-backend是当前 rendezvous 通信的后端方法,推荐使用 c10d。
  • --rdzv-endpoint是 rendezvous 通信后端的地址,即 master 节点地址和端口。
  • --nnodes是当前参与分布式训练的节点数量,即当前 role 的节点数量。可使用环境变量VC_WORKER_NUM获取当前 worker role 节点数量。
  • --nproc_per_node=8是每个节点使用的 GPU 数量。
  • YOUR_TRAINING_SCRIPT.py是您的启动脚本。
  • <script args>是启动脚本接收的参数。

更多环境变量介绍可参考创建分布式训练模块中的“平台预置环境变量“。

使用 torch.distributed.launch 启动训练​

以“一个 worker role 共2个节点,每个节点 8 个 GPU” 为例,输入启动命令如下:

# 取 worker 节点组中第 1 个节点作为 master 节点
export MASTER=`echo ${VC_WORKER_HOSTS} | awk -F , '{print $1}'`

python -m torch.distributed.launch \
--nnodes=2 \
--nproc-per-node=8 \
--node-rank=${VC_TASK_INDEX} \
--master-addr=${MASTER} \
--master-port=29500 \
YOUR_TRAINING_SCRIPT.py \
<script args>

其中:

  • --nnodes是当前参与分布式训练的节点数量,即当前 role 的节点数量。可使用环境变量VC_WORKER_NUM获取当前 worker role 节点数量。
  • --nproc_per_node是每个节点使用的 GPU 数量。
  • --node-rank是节点序号,环境变量 VC_TASK_INDEX 是当前节点在当前 role(节点组)中的序号,从0开始编号。
  • --master-addr是 master 节点 IP 或域名。
  • --master-port是 master 节点通信端口,任意一个未占用端口即可。
  • YOUR_TRAINING_SCRIPT.py是您的启动脚本。
  • <script args>是启动脚本接收的参数。

更多环境变量介绍可参考创建分布式训练模块中的“平台预置环境变量“。