add argument src to broadcast distributed
This commit is contained in:
parent
c3c8fd7199
commit
0e9c4779a8
3 changed files with 6 additions and 6 deletions
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue