最近想上手试试用PyTorch 2.0自带的DDP(Distributed DataParallel)在两张4090上跑一个7B参数的LLaMA风格模型。我原以为开多卡至少能快个1.5倍,结果实际测试下来,单卡跑一个batch大概1.2秒,双卡DDP反而要1.5秒,甚至偶尔还会卡到2秒以上。
楼主
3天前
用PyTorch2.0训7B模型,DDP比单卡还慢,是我哪里姿势不对吗?
请 登录 后发表回复
全部回复
共 43 条
2楼
13小时前
大概率是allreduce通信开销吃掉了双卡收益,7B模型每步同步的数据量太大了,试试调大batch size或者用梯度累积压一下。
3楼
12小时前
我猜是不是batch size太小了,DDP通信开销没摊平,试试加大batch看看。
4楼
11小时前
这情况我也遇到过,DDP在小batch下通信开销占比太高了,特别是7B这种大模型,单卡batch_size如果不够大,多卡反而会因为梯度同步拖慢速度。建议你试试增大每卡的batch_size,或者开启PyTorch 2.0的compile模式,有时候能缓解通信瓶颈。另外检查下NVLink是否生效,两张4090如果走PCIE通信,延迟会比较明显。