使用SMDDP时遇SageMaker Data Parallel错误:DDP不支持broadcast_buffers
修复SageMaker Data Parallel的broadcast_buffers不支持错误
解决方案
初始化SMDDP的DDP时,显式设置broadcast_buffers=False参数即可解决该错误。
修改后的代码片段
... import smdistributed.dataparallel.torch.distributed as dist from smdistributed.dataparallel.torch.parallel.distributed import DistributedDataParallel as DDP ... if __name__ == "__main__": ... model = models.segmentation.deeplabv3_mobilenet_v3_large( pretrained=False, progress=False, num_classes=args.classes) # 显式关闭broadcast_buffers model = DDP(model, broadcast_buffers=False) model.train() amp = True CE = torch.nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=args.lr, momentum=args.momentum) scaler = GradScaler(enabled=amp) ...
原因说明
SMDDP(SageMaker Data Parallel)的DDP实现与PyTorch原生DDP存在差异,它不支持broadcast_buffers=True的默认配置,必须手动将该参数设置为False才能兼容运行。
内容的提问来源于stack exchange,提问作者Arvs
相关产品推荐
相关产品推荐

