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):