test broadcast and all reduce
This commit is contained in:
parent
55db13836b
commit
94b3582a46
5 changed files with 17 additions and 14 deletions
|
|
@ -46,17 +46,13 @@ void broadcast_tensor(Tensor* tensor) {
|
|||
cudaStreamDestroy(stream);
|
||||
}
|
||||
|
||||
void allreduce_mean_tensor(Tensor* tensor) {
|
||||
void allreduce_sum_tensor(Tensor* tensor) {
|
||||
cudaStream_t stream;
|
||||
cudaStreamCreate(&stream);
|
||||
|
||||
// Perform NCCL AllReduce operation to calculate the sum of all tensors across all processes
|
||||
NCCL_CHECK(ncclAllReduce(tensor->data, tensor->data, tensor->size, ncclFloat, ncclSum, nccl_comm, stream));
|
||||
|
||||
// Divide the sum by the world size to get the mean
|
||||
float scale = 1.0f / world_size;
|
||||
scalar_mul_tensor_cuda(tensor, scale, tensor->data);
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaStreamDestroy(stream);
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue