PyTorch中DistributedDataParallel使用笔记

1. 基本概念

在使用DistributedDataParallel时有一些概率必须掌握

多机多卡 含义
world_size 代表有几台机器,可以理解为几台服务器
rank 第几台机器,即第几个服务器
local_rank 某台机器中的第几块GPU
单机多卡 含义
world_size 代表机器一共有几块GPU
rank 第几块GPU
local_rank 第几块GPU,与rank相同

2. 使用方法

2.1. 修改主函数

在运行的时候,DistributedDataParallel会往你的程序中加入一个参数local_rank,所以要现在你的代码中解析这个参数,如:

python 复制代码
parser.add_argument("--local_rank", type=int, default=1, help="number of cpu threads to use during batch generation")

2.2. 初始化

python 复制代码
torch.distributed.init_process_group(backend="nccl")

os.environ["CUDA_VISIBLE_DEVICES"] = "0, 1, 2"  # 有几块GPU写多少

2.3. 设定device

python 复制代码
local_rank = torch.distributed.get_rank()
torch.cuda.set_device(local_rank)
global device
device = torch.device("cuda", local_rank)

我没用arg.local_rank,新定义了一个local_rank变量,是因为我更信任distributed.get_rank()这个函数

这里用torch.device来写,并且加了global,是因为后面模型和数据都要用到这个device,不会出错

2.4. 模型加载到多gpu

python 复制代码
model.to(device)  # 这句不能少,最好不要用model.cuda()
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank, find_unused_parameters=True)  # 这句加载到多GPU上

2.5. 数据加载到gpu

python 复制代码
数据.to(device)

2.6. 启动

python 复制代码
torchrun --nproc_per_node=4 --rdzv_endpoint=localhost:12345 train_cylinder_asym.py

参考文献

Pytorch并行计算(二): DistributedDataParallel介绍_dist.barrier_harry_tea的博客-CSDN博客

DistributedDataParallel多GPU分布式训练全过程总结 跟着做90%成功_BRiAq的博客-CSDN博客

相关推荐
AI人工智能+1 天前
学位证书识别技术通过计算机视觉、自然语言处理和大模型推理相结合,实现了高精度、高效率的证书数字化处理
深度学习·计算机视觉·自然语言处理·ocr·学位证书识别
ZHOU_WUYI1 天前
5. fastwam 模型 video expert pre dit过程
pytorch·python·世界动作模型
过期的秋刀鱼!1 天前
机器学习开发的迭代循环
人工智能·python·深度学习·神经网络·机器学习
向哆哆1 天前
水稻病害检测数据集分享(适用于YOLO系列深度学习分类检测任务)
深度学习·yolo·分类
林泽毅1 天前
PyTRIO快速入门实战篇(二):用 GRPO 提升 GSM8K 数学推理准确率
人工智能·深度学习·机器学习·llm·强化学习
AI人工智能+2 天前
一种基于深度学习技术的高精度医疗机构执业许可证识别系统,构建了一套基于深度神经网络的端到端智能识别系统,为医疗行业提
深度学习·ocr·医疗机构执业许可证识别
硅谷秋水2 天前
EgoSteer:一种基于第一人称视角视频、实现可控灵巧操作的全栈系统
深度学习·机器学习·语言模型·机器人·音视频
AI街潜水的八角2 天前
基于YOLO26交通标志检测系统1:交通标志检测数据集说明(含下载链接)
深度学习·神经网络
aiblog2 天前
深度学习中“Transformer”怎么翻译为中文?
人工智能·深度学习·transformer
AndrewHZ2 天前
【LLM技术全景】阶段总结:技术原理篇核心知识回顾
人工智能·深度学习·算法·语言模型·大模型·llm·芯片开发