willychan21/ParallelKernelBench_Problems
ParallelKernelBench (benchmark) Reference problems for ParallelKernelBench: a benchmark for LLM-generated multi-GPU CUDA kernels. This dataset contains 87 reference implementations in reference/ and the input tensor specification in utils/input_output_tensors.py. Files Path Description data/problems.parquet One row per problem (tabular access) reference/*.py Reference solution() implementations utils/input_output_tensors.py Input/output tensor… See the full description on the dataset page: https://huggingface.co/datasets/willychan21/ParallelKernelBench_Problems.
0123
1import torch2import torch.distributed as dist3from typing import Tuple4 5def rotate_half(x: torch.Tensor) -> torch.Tensor:6 half_dim = x.shape[-1] // 27 x1, x2 = x[..., :half_dim], x[..., half_dim:]8 return torch.cat((-x2, x1), dim=-1)9 10def solution(11 q_local: torch.Tensor, 12 k_local: torch.Tensor, 13 cos_local: torch.Tensor, 14 sin_local: torch.Tensor15) -> Tuple[torch.Tensor, torch.Tensor]:16 # Reshape cos and sin to broadcast with q and k's head dimension (dim=2)17 cos = cos_local.unsqueeze(2)18 sin = sin_local.unsqueeze(2)19 20 q_embed_local = (q_local * cos) + (rotate_half(q_local) * sin)21 k_embed_local = (k_local * cos) + (rotate_half(k_local) * sin)22 23 if not dist.is_initialized():24 return q_embed_local, k_embed_local25 26 world_size = dist.get_world_size()27 28 q_gather_list = [torch.empty_like(q_embed_local) for _ in range(world_size)]29 k_gather_list = [torch.empty_like(k_embed_local) for _ in range(world_size)]30 31 dist.all_gather(q_gather_list, q_embed_local.contiguous())32 dist.all_gather(k_gather_list, k_embed_local.contiguous())33 34 q_embed_global = torch.cat(q_gather_list, dim=1)35 k_embed_global = torch.cat(k_gather_list, dim=1)36 37 return q_embed_global, k_embed_global38 