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
1from __future__ import annotations2 3import torch4import torch.distributed as dist5from torch import Tensor6 7 8@torch.no_grad()9def solution(10 rs_input_1d: Tensor,11 gamma: Tensor,12 eps: float,13) -> Tensor:14 world_size = dist.get_world_size()15 n = rs_input_1d.numel()16 chunk = n // world_size17 18 hidden = gamma.numel()19 assert chunk % hidden == 0, f"chunk ({chunk}) must divide hidden ({hidden})"20 rows = chunk // hidden21 22 out_flat = torch.empty(chunk, dtype=rs_input_1d.dtype, device=rs_input_1d.device)23 dist.reduce_scatter_tensor(out_flat, rs_input_1d.contiguous(), op=dist.ReduceOp.SUM)24 out_flat.div_(world_size)25 26 x = out_flat.view(rows, hidden).float()27 gn = gamma.float()28 rms = torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True).add(eps))29 y = x * rms * gn30 return y.to(dtype=rs_input_1d.dtype)31 