feat: use D2D instead of H2H in pp (#7673)

Co-authored-by: alpha-baby <fujianhao1997@qq.com>
This commit is contained in:
TianyuZhang1214
2025-07-04 01:58:50 +08:00
committed by GitHub
parent 264dc6e744
commit 0099172327
3 changed files with 45 additions and 22 deletions

View File

@@ -1000,36 +1000,48 @@ def point_to_point_pyobj(
src: int = 0,
dst: int = 1,
):
"""Send data from src to dst in group."""
"""Send data from src to dst in group using DeviceToDevice communication."""
if rank == src:
if len(data) == 0:
tensor_size = torch.tensor([0], dtype=torch.long)
tensor_size = torch.tensor(
[0], dtype=torch.long, device=torch.cuda.current_device()
)
dist.send(tensor_size, dst=dst, group=group)
else:
serialized_data = pickle.dumps(data)
size = len(serialized_data)
tensor_data = torch.ByteTensor(
np.frombuffer(serialized_data, dtype=np.uint8)
).cuda(
device=torch.cuda.current_device()
) # Move to GPU
tensor_size = torch.tensor(
[size], dtype=torch.long, device=torch.cuda.current_device()
)
tensor_size = torch.tensor([size], dtype=torch.long)
dist.send(tensor_size, dst=dst, group=group)
dist.send(tensor_data, dst=dst, group=group)
return data
elif rank == dst:
tensor_size = torch.tensor([0], dtype=torch.long)
tensor_size = torch.tensor(
[0], dtype=torch.long, device=torch.cuda.current_device()
)
dist.recv(tensor_size, src=src, group=group)
size = tensor_size.item()
if size == 0:
return []
tensor_data = torch.empty(size, dtype=torch.uint8)
tensor_data = torch.empty(
size, dtype=torch.uint8, device=torch.cuda.current_device()
)
dist.recv(tensor_data, src=src, group=group)
serialized_data = bytes(tensor_data.cpu().numpy())
serialized_data = bytes(
tensor_data.cpu().numpy()
) # Move back to host for deserialization
data = pickle.loads(serialized_data)
return data