add argument src to broadcast distributed

This commit is contained in:
lucasdelimanogueira 2024-05-25 09:38:14 -03:00
parent c3c8fd7199
commit 0e9c4779a8
3 changed files with 6 additions and 6 deletions

View file

@ -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);
}

View file

@ -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);
}

View file

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