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
45_reducescatter_fused_rmsnorm.py31 linesDownload Raw Back to reference
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