最近在折腾 MCP 框架,想在多卡环境里跑一个简单的 ResNet 训练。我按官方文档搭了 DDP,但发现 loss 下降曲线和单卡跑出来的差挺多,怀疑是梯度同步没弄对。具体场景是:4 张 V100,batch size 每卡 32,用了 DistributedSampler,但 log 里每步的 loss 数值波动很大,不像单卡那么平滑。我检查了 torch.distributed.all_reduce 的调用位置,但不确定是不是因为有的 layer 忘了加 sync。另外,MCP 的 checkpoint 保存方式也和原生 DDP 不太一样,有点懵。有没有踩过类似坑的大神,能讲讲常见错误或者调试思路?感谢!
楼主
4小时前
MCP 用 PyTorch 做分布式训练,梯度同步总感觉不对,求指点
请 登录 后发表回复
全部回复
共 4 条
2楼
3小时前
我之前也遇到过类似的问题,后来发现是DistributedSampler的shuffle没处理好,导致每个epoch的数据分配不一致,loss自然跳得厉害。另外你提到的all_reduce,建议检查下是不是某些自定义layer或者BN层没自动同步,手动hook一下会稳很多。MCP的checkpoint我建议直接基于state_dict自己写保存逻辑,别完全依赖框架自带的方法,这样后续加载也更灵活。
3楼
3小时前
检查下是不是model在DDP包裹前就调用了all_reduce,或者试试把BN换成SyncBN看看。
4楼
1小时前
我之前也遇到过类似情况,后来发现是DistributedSampler在dataloader里没设shuffle=True,导致每个epoch的样本顺序不一样,loss自然会跳得厉害。另外建议检查一下model.parameters()里有没有哪些层被设了requires_grad=False,可能会让all_reduce漏掉梯度。关于MCP的checkpoint,我记得它默认是只保存rank0的状态,得手动把整个模型的state_dict汇总一下才行,你可以看看官方示例里是怎么处理的。
5楼
43分钟前
检查下是不是用了model.no_sync()但没在backward前手动同步,MCP的hook写法容易漏这个。