集合通信动画演示

min

本文通过交互动画和 PyTorch 示例,直观展示常见集合通信操作中数据如何在 GPU 之间流动。

交互演示

选择通信原语和 GPU 数量后,即可连续播放或逐轮观察。

AllReduce 全规约 · Ring 初始状态

你的浏览器不支持 Canvas 动画。
准备就绪

查看逐轮通信记录

    PyTorch CPU 代码示例

    下面的示例使用 PyTorch Distributed 的 Gloo 后端,在单机启动 4 个进程;保存代码后运行 python 文件名.py 即可。

    Broadcast:broadcast_cpu_demo.py
    # 运行方式: python broadcast_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29500",
        )
    
        if rank == 0:
            tensor = torch.tensor([1.0, 2.0, 3.0])
        else:
            tensor = torch.zeros(3)
    
        print(f"[rank {rank}] before: {tensor.tolist()}")
        dist.broadcast(tensor, src=0)
        print(f"[rank {rank}] after : {tensor.tolist()}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    Scatter:scatter_cpu_demo.py
    # 运行方式: python scatter_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29501",
        )
    
        output_tensor = torch.zeros(2)
    
        if rank == 0:
            full_data = torch.arange(2 * world_size, dtype=torch.float32)
            scatter_list = list(full_data.chunk(world_size))
            print(f"[rank 0] full_data = {full_data.tolist()}")
        else:
            scatter_list = None
    
        dist.scatter(output_tensor, scatter_list=scatter_list, src=0)
        print(f"[rank {rank}] received: {output_tensor.tolist()}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    Gather:gather_cpu_demo.py
    # 运行方式: python gather_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29502",
        )
    
        my_chunk = torch.full((2,), fill_value=float(rank))
        print(f"[rank {rank}] my_chunk = {my_chunk.tolist()}")
    
        gather_list = [torch.zeros(2) for _ in range(world_size)] if rank == 0 else None
    
        dist.gather(my_chunk, gather_list=gather_list, dst=0)
    
        if rank == 0:
            print(f"[rank 0] gathered: {[t.tolist() for t in gather_list]}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    Reduce:reduce_cpu_demo.py
    # 运行方式: python reduce_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29503",
        )
    
        tensor = torch.full((4,), fill_value=float(rank + 1))
        print(f"[rank {rank}] before: {tensor.tolist()}")
    
        # 只保证 dst 拿到正确的规约结果
        dist.reduce(tensor, dst=0, op=dist.ReduceOp.SUM)
    
        # 非 dst 的 tensor 内容未定义
        print(f"[rank {rank}] after : {tensor.tolist()}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    AllGather:allgather_cpu_demo.py
    # 运行方式: python allgather_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29504",
        )
    
        my_chunk = torch.full((2,), fill_value=float(rank))
        print(f"[rank {rank}] my_chunk = {my_chunk.tolist()}")
    
        # 所有 rank 都要准备长度为 world_size 的接收缓冲区列表
        gather_list = [torch.zeros(2) for _ in range(world_size)]
    
        dist.all_gather(gather_list, my_chunk)
    
        print(f"[rank {rank}] gathered: {[t.tolist() for t in gather_list]}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    ReduceScatter:reducescatter_cpu_demo.py
    # 运行方式: python reducescatter_cpu_demo.py
    #
    # Gloo 不支持原生 dist.reduce_scatter(),这里用 world_size 次
    # dist.reduce(dst=j) 组合出等价效果。
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29505",
        )
    
        # input_list[j] 是当前 rank 要贡献给 rank j 的输入
        input_list = [
            torch.full((2,), fill_value=float(rank * 10 + j))
            for j in range(world_size)
        ]
        print(f"[rank {rank}] input_list = {[t.tolist() for t in input_list]}")
    
        output = torch.zeros(2)
        for j in range(world_size):
            part = input_list[j].clone()
            dist.reduce(part, dst=j, op=dist.ReduceOp.SUM)
            if rank == j:
                output = part
    
        print(f"[rank {rank}] reduce_scatter result: {output.tolist()}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    AllReduce:allreduce_cpu_demo.py
    # 运行方式: python allreduce_cpu_demo.py
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29506",
        )
    
        tensor = torch.full((4,), fill_value=float(rank + 1))
        print(f"[rank {rank}] before: {tensor.tolist()}")
    
        # 原地操作,所有 rank 都会拿到相同的求和结果
        dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
    
        print(f"[rank {rank}] after : {tensor.tolist()}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)
    All-to-All:alltoall_cpu_demo.py
    # 运行方式: python alltoall_cpu_demo.py
    #
    # Gloo 不支持列表形式的 dist.all_to_all(),改用等价接口
    # dist.all_to_all_single()。
    import torch
    import torch.distributed as dist
    import torch.multiprocessing as mp
    
    
    def worker(rank, world_size):
        dist.init_process_group(
            backend="gloo", rank=rank, world_size=world_size,
            init_method="tcp://127.0.0.1:29507",
        )
    
        chunk_size = 2
        # 第 j 段是当前 rank 要发给 rank j 的个性化数据
        input_tensor = torch.cat([
            torch.full((chunk_size,), fill_value=float(rank * 10 + j))
            for j in range(world_size)
        ])
        print(f"[rank {rank}] input_tensor = {input_tensor.tolist()}")
    
        output_tensor = torch.zeros_like(input_tensor)
        dist.all_to_all_single(output_tensor, input_tensor)
    
        # output_tensor 第 j 段就是 rank j 发来的数据
        received = list(output_tensor.chunk(world_size))
        print(f"[rank {rank}] received: {[t.tolist() for t in received]}")
    
        dist.destroy_process_group()
    
    
    if __name__ == "__main__":
        world_size = 4
        mp.spawn(worker, args=(world_size,), nprocs=world_size, join=True)