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
0likes146downloads
72_hyena_forward_cp.py142 linesDownload Raw Back to reference
1from typing import Optional2 3import torch4import torch.distributed as dist5 6 7def _zigzag_indices(num_chunks: int, device: torch.device) -> torch.Tensor:8    half = (num_chunks + 1) // 29    left = torch.arange(half, device=device)10    right = torch.arange(num_chunks - 1, half - 1, -1, device=device)11    indices = torch.empty(num_chunks, dtype=torch.long, device=device)12    indices[0::2] = left13    indices[1::2] = right14    return indices15 16 17def _inverse_zigzag_indices(num_chunks: int, device: torch.device) -> torch.Tensor:18    half = num_chunks // 219    left = torch.arange(half, device=device)20    right = torch.arange(num_chunks - 1, half - 1, -1, device=device)21    indices = torch.empty(num_chunks, dtype=torch.long, device=device)22    indices[0::2] = left23    indices[1::2] = right24    return torch.argsort(indices)25 26 27def _a2a_split_to_full(28    x: torch.Tensor,29    group: dist.ProcessGroup,30    with_zigzag_splitting: bool,31) -> torch.Tensor:32    world_size = dist.get_world_size(group=group)33    batch, global_channels, local_seq = x.shape34    local_channels = global_channels // world_size35    seq_len = local_seq * world_size36 37    send = (38        x.reshape(batch, world_size, local_channels, local_seq)39        .permute(1, 0, 2, 3)40        .contiguous()41    )42    recv = torch.empty_like(send)43    dist.all_to_all_single(recv, send, group=group)44    out = (45        recv.permute(1, 2, 0, 3)46        .reshape(batch, local_channels, seq_len)47        .contiguous()48    )49 50    if with_zigzag_splitting:51        num_chunks = 2 * world_size52        index = _inverse_zigzag_indices(num_chunks, out.device)53        out = (54            out.reshape(batch, local_channels, num_chunks, seq_len // num_chunks)55            .index_select(dim=2, index=index)56            .reshape(batch, local_channels, seq_len)57        )58    return out59 60 61def _a2a_full_to_split(62    x: torch.Tensor,63    group: dist.ProcessGroup,64    with_zigzag_splitting: bool,65) -> torch.Tensor:66    world_size = dist.get_world_size(group=group)67    batch, local_channels, seq_len = x.shape68    local_seq = seq_len // world_size69 70    if with_zigzag_splitting:71        num_chunks = 2 * world_size72        index = _zigzag_indices(num_chunks, x.device)73        x = (74            x.reshape(batch, local_channels, num_chunks, seq_len // num_chunks)75            .index_select(dim=2, index=index)76            .reshape(batch, local_channels, seq_len)77        )78 79    send = (80        x.reshape(batch, local_channels, world_size, local_seq)81        .permute(2, 0, 1, 3)82        .contiguous()83    )84    recv = torch.empty_like(send)85    dist.all_to_all_single(recv, send, group=group)86    return (87        recv.permute(1, 0, 2, 3)88        .reshape(batch, world_size * local_channels, local_seq)89        .contiguous()90    )91 92 93def _fftconv_ref(94    u: torch.Tensor,95    kernel: torch.Tensor,96    bias: torch.Tensor,97) -> torch.Tensor:98    seq_len = u.shape[-1]99    fft_size = 2 * seq_len100    u_float = u.float()101    kernel_float = kernel.float()102 103    kernel_f = torch.fft.rfft(kernel_float, n=fft_size) / fft_size104    u_f = torch.fft.rfft(u_float, n=fft_size)105    y = torch.fft.irfft(106        u_f * kernel_f.unsqueeze(0), n=fft_size, norm="forward"107    )[..., :seq_len]108    y = y + u_float * bias.float().unsqueeze(-1)109    return y.to(dtype=u.dtype)110 111 112@torch.no_grad()113def solution(114    x1_seq: torch.Tensor,115    x2_seq: torch.Tensor,116    v_seq: torch.Tensor,117    h: torch.Tensor,118    conv_bias: torch.Tensor,119    num_groups: int,120    group_dim: int,121    group: Optional[dist.ProcessGroup] = None,122    with_zigzag_splitting: bool = True,123) -> torch.Tensor:124    group = group or dist.group.WORLD125    world_size = dist.get_world_size(group=group)126    rank = dist.get_rank(group=group)127 128    x1 = _a2a_split_to_full(x1_seq, group, with_zigzag_splitting)129    x2 = _a2a_split_to_full(x2_seq, group, with_zigzag_splitting)130    v = _a2a_split_to_full(v_seq, group, with_zigzag_splitting)131 132    local_channels = x1.shape[1]133    local_groups = num_groups // world_size134    h_local = h[rank * local_groups : (rank + 1) * local_groups]135    h_local = h_local.repeat_interleave(group_dim, dim=0)136    bias_local = conv_bias[rank * local_channels : (rank + 1) * local_channels]137 138    z = x2 * v139    z = _fftconv_ref(z, h_local, bias_local)140    z = x1 * z141    z_seq = _a2a_full_to_split(z, group, with_zigzag_splitting)142    return z_seq.transpose(1, 2).contiguous()