Datasets:
Download reference/11_allgather_gemm_AT.py from togethercomputer/ParallelKernelBench_Problems: direct link, hf CLI and curl.
- Browser
- Download file 525 Bytes
-
https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/main/reference/11_allgather_gemm_AT.py
- Command line
-
hf download hf://datasets/togethercomputer/ParallelKernelBench_Problems/reference/11_allgather_gemm_AT.py
-
curl -L -o 11_allgather_gemm_AT.py https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/main/reference/11_allgather_gemm_AT.py
525 Bytes
| import torch | |
| import torch.distributed as dist | |
| def solution(A_local: torch.Tensor, B: torch.Tensor) -> torch.Tensor: | |
| world_size = dist.get_world_size() | |
| M, K_local = A_local.shape | |
| K = world_size * K_local | |
| A_local_t = A_local.transpose(0, 1).contiguous() | |
| A_t_buf = A_local_t.new_empty((world_size, K_local, M)) | |
| dist.all_gather_into_tensor(A_t_buf, A_local_t) | |
| A_global_t = A_t_buf.reshape(K, M) | |
| C_t = torch.matmul(B.transpose(0, 1), A_global_t) | |
| return C_t.transpose(0, 1) | |