From 0e9c4779a8c689f6e77c9bb9f7ee2f0c9ec86266 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Sat, 25 May 2024 09:38:14 -0300 Subject: [PATCH] add argument src to broadcast distributed --- norch/csrc/distributed.cpp | 4 ++-- norch/csrc/distributed.h | 2 +- norch/distributed/distributed.py | 6 +++--- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/norch/csrc/distributed.cpp b/norch/csrc/distributed.cpp index 5239eaa..bb753b2 100644 --- a/norch/csrc/distributed.cpp +++ b/norch/csrc/distributed.cpp @@ -35,12 +35,12 @@ void init_process_group(int env_rank, int env_world_size) { } -void broadcast_tensor(Tensor* tensor) { +void broadcast_tensor(Tensor* tensor, int src) { cudaStream_t stream; cudaStreamCreate(&stream); - NCCL_CHECK(ncclBroadcast(tensor->data, tensor->data, tensor->size * sizeof(float), ncclFloat, 0, nccl_comm, stream)); + NCCL_CHECK(ncclBroadcast(tensor->data, tensor->data, tensor->size * sizeof(float), ncclFloat, src, nccl_comm, stream)); cudaStreamSynchronize(stream); cudaStreamDestroy(stream); } diff --git a/norch/csrc/distributed.h b/norch/csrc/distributed.h index 3a94983..53e318e 100644 --- a/norch/csrc/distributed.h +++ b/norch/csrc/distributed.h @@ -22,7 +22,7 @@ extern "C" { void init_process_group(int rank, int world_size); - void broadcast_tensor(Tensor* tensor); + void broadcast_tensor(Tensor* tensor, int src); void allreduce_sum_tensor(Tensor* tensor); void allreduce_mean_tensor(Tensor* tensor); } diff --git a/norch/distributed/distributed.py b/norch/distributed/distributed.py index f84a7fa..bd87bc3 100644 --- a/norch/distributed/distributed.py +++ b/norch/distributed/distributed.py @@ -9,12 +9,12 @@ def init_process_group(rank, world_size, backend='nccl'): Tensor._C.init_process_group(rank, world_size) -def broadcast_tensor(tensor): +def broadcast_tensor(tensor, src=0): - Tensor._C.broadcast_tensor.argtypes = [ctypes.POINTER(CTensor)] + Tensor._C.broadcast_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.c_int] Tensor._C.broadcast_tensor.restype = None - Tensor._C.broadcast_tensor(tensor.tensor) + Tensor._C.broadcast_tensor(tensor.tensor, src) def allreduce_sum_tensor(tensor):