You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 07:45:29