Datasets:
Download reference/46_gemv_decode.py from togethercomputer/ParallelKernelBench_Problems: direct link, hf CLI and curl.
- Browser
- Download file 535 Bytes
-
https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/main/reference/46_gemv_decode.py
- Command line
-
hf download hf://datasets/togethercomputer/ParallelKernelBench_Problems/reference/46_gemv_decode.py
-
curl -L -o 46_gemv_decode.py https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems/resolve/main/reference/46_gemv_decode.py
535 Bytes
| import torch | |
| import torch.distributed as dist | |
| def solution( | |
| hidden_states: torch.Tensor, | |
| weight_shard: torch.Tensor, | |
| bias_shard: torch.Tensor, | |
| ) -> torch.Tensor: | |
| world_size = dist.get_world_size() | |
| local_logits = torch.matmul(hidden_states, weight_shard.t()) | |
| local_logits = local_logits + bias_shard | |
| gathered = [torch.empty_like(local_logits) for _ in range(world_size)] | |
| dist.all_gather(gathered, local_logits.contiguous()) | |
| logits = torch.cat(gathered, dim=1) | |
| return logits | |