Team Ai
Datasetpublic

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.

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
0likes123downloads
17_rope_allgather.py38 linesDownload Raw Back to reference
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