集合通信动画演示
min
本文通过交互动画和 PyTorch 示例,直观展示常见集合通信操作中数据如何在 GPU 之间流动。
交互演示
选择通信原语和 GPU 数量后,即可连续播放或逐轮观察。
AllReduce 全规约 · Ring 初始状态
准备就绪
查看逐轮通信记录
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)