PyTorch分布式all_gather数据碰撞:写入文件可修复该问题
PyTorch分布式all_gather索引碰撞诡异问题
问题现象
- 分布式训练场景下,各进程计算误差后返回附带
float类型索引的结果 - 使用
torch.distributed.all_gather收集不同rank的数据时,出现索引碰撞:例如100个样本、4块GPU的场景,最终索引集合长度仅为25,远低于预期的100 - 前置验证:将各rank聚合前的数据写入文件,可确认不同rank的索引完全无重叠
- 诡异的临时“修复”现象:
- 聚合后将数据写入文件,索引碰撞问题消失(
len(set(data.numpy())) == 100) - 注释掉该文件写入代码,问题立刻重现(
len(set(data.numpy())) == 25) - 仅打印聚合后的结果也能达到同样的“修复”效果,但对聚合结果进行排序操作无法解决问题
- 聚合后将数据写入文件,索引碰撞问题消失(
疑似原因猜测
目前怀疑是分布式操作后的数据内存可见性问题,类似IO操作需要flush来确保数据同步的情况,但在torch.distributed.all_gather的官方文档中未找到相关说明。
复现代码
# setup_distributed_stuff() rank = torch.distributed.get_rank() world_size = torch.distributed.get_world_size() # 分布式计算返回的数据,各rank间无重叠 data = torch.arange( 0 + (rank * 100 // world_size), (rank + 1) * 100 // world_size, ) # 写入文件可确认各rank的data无重叠 # 从所有rank收集数据 if world_size > 1: all_data = [torch.zeros_like(data) for _ in range(world_size)] torch.distributed.all_gather(all_data, data) data = torch.cat(all_data, dim=0) # 将data写入文件调试时,问题消失:len(set(data.numpy())) == 100 # 注释该代码后,收集的数据出现碰撞:len(set(data.numpy())) == 100 // world_size with open("debug_data.pt", "wb") as _file: torch.save(data, _file) # 仅打印索引也能得到同样效果 logger.info( "Gathered result indices: {}...{}".format( data[:10, -1], data[-10:, -1] ) ) # 但对索引排序无济于事 data = data[data[:, -1].argsort(dim=0)] if rank == 0: # do_something(data)
内容的提问来源于stack exchange,提问作者aweeeezy
相关产品推荐
相关产品推荐

