最近想上手试试用PyTorch 2.0自带的DDP(Distributed DataParallel)在两张4090上跑一个7B参数的LLaMA风格模型。我原以为开多卡至少能快个1.5倍,结果实际测试下来,单卡跑一个batch大概1.2秒,双卡DDP反而要1.5秒,甚至偶尔还会卡到2秒以上。
用PyTorch2.0训7B模型,DDP比单卡还慢,是我哪里姿势不对吗?
全部回复
共 186 条7B模型DDP通信开销很大,两张4090没NVLink,梯度all-reduce走PCIe肯定拖后腿。试试开gradient accumulation或者用FSDP。
7B模型DDP变慢挺常见的,尤其是小batch下梯度all-reduce的通信开销占比太高,两张4090走PCIe可能直接把收益吃掉了。你试试把batch size调大一点,或者开gradient accumulation,让计算/通信比上来。另外确认下有没有开find_unused_parameters,这玩意会拖慢不少。单卡1.2秒本来就没到瓶颈,双卡想快1.5倍有点理想化了。
7B模型DDP反而变慢挺常见的,八成是通信开销把并行收益吃掉了。两张4090之间如果是PCIe走数据,梯度all-reduce那一下延迟很高,模型越大越明显。你可以先试试开gradient accumulation模拟更大batch,或者检查下是不是没设bucket_cap_mb,默认25MB对小梯度同步不太友好。另外确认下数据加载有没有成瓶颈,有时候卡在dataloader上,多卡反而抢IO更凶。
7B模型在两张4090上DDP变慢挺常见的,不一定是姿势问题。4090没有NVLink,卡间通信走PCIe,梯度all-reduce那一下开销很实在,模型越大越明显。你单卡1.2秒一个batch,说明计算本身没吃满,DDP的通信反而成了纯增量。可以先用torchrun加NCCL_DEBUG=INFO看看是不是走了PCIe而不是P2P。另外确认下有没有开gradient_as_bucket_view和static_graph,PyTorch 2.0里这两个对DDP的overlap帮助不小。还有个容易忽略的点,DataLoader的worker数和batch size如果没随卡数调整,双卡可能一直在等数据。建议先把batch size翻倍再测,不然单卡双卡比的根本不是同一件事。
7B模型DDP通信开销本来就大,两张4090没NVLink,梯度同步走PCIe能不慢吗?
两张4090跑7B还变慢,大概率是通信开销把收益吃掉了。4090没有NVLink,走PCIe传梯度本来就不便宜,模型越大同步的参数量越吓人。你可以先看看是不是batch太小,单卡利用率都没跑满,DDP反而多了一层all-reduce。另外试试开gradient accumulation或者把bucket_cap_mb调大点,有时候能救回来一些。