Datasets:
Download reference/13_gemm_allreduce.py from togethercomputer/ParallelKernelBench_Problems: direct link, hf CLI and curl.
- Browser
- Download file 486 Bytes
-
https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/3ea6e263b29754fb639df0e25530c387f732fbcb/reference/13_gemm_allreduce.py
- Command line
-
hf download hf://datasets/togethercomputer/ParallelKernelBench_Problems@3ea6e263b29754fb639df0e25530c387f732fbcb/reference/13_gemm_allreduce.py
-
curl -L -o 13_gemm_allreduce.py https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/3ea6e263b29754fb639df0e25530c387f732fbcb/reference/13_gemm_allreduce.py
486 Bytes
| import torch | |
| import torch.distributed as dist | |
| def solution( | |
| A_local: torch.Tensor, | |
| B_local: torch.Tensor, | |
| ) -> torch.Tensor: | |
| rank = dist.get_rank() | |
| world_size = dist.get_world_size() | |
| M, K = A_local.shape | |
| K_B, N = B_local.shape | |
| A_local = A_local.contiguous() | |
| B_local = B_local.contiguous() | |
| C_local = torch.matmul(A_local, B_local) | |
| C = C_local.clone() | |
| dist.all_reduce(C, op=dist.ReduceOp.SUM) | |
| return C | |