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 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()