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.
0263
1from typing import Dict, List, Optional, Tuple2 3import torch4import torch.distributed as dist5 6 7def _sum_by_splits(values: List[int], splits: List[int]) -> List[int]:8 out: List[int] = []9 offset = 010 for split in splits:11 out.append(sum(values[offset : offset + split]))12 offset += split13 return out14 15 16def _lengths_per_key(lengths: torch.Tensor, stride_per_key: List[int]) -> List[int]:17 out: List[int] = []18 offset = 019 for stride in stride_per_key:20 out.append(int(lengths[offset : offset + stride].sum().item()))21 offset += stride22 return out23 24 25def _get_recat(26 local_split: int,27 num_splits: int,28 stagger: int = 1,29 device: Optional[torch.device] = None,30 batch_size_per_rank: Optional[List[int]] = None,31) -> Optional[torch.Tensor]:32 if local_split == 0:33 return None34 35 feature_order = [36 x + num_splits // stagger * y37 for x in range(num_splits // stagger)38 for y in range(stagger)39 ]40 if batch_size_per_rank is None:41 recat = [42 feature_idx + rank_idx * local_split43 for feature_idx in range(local_split)44 for rank_idx in feature_order45 ]46 else:47 rank_offsets = [0]48 for batch_size in batch_size_per_rank[:-1]:49 rank_offsets.append(rank_offsets[-1] + local_split * batch_size)50 recat = [51 rank_offsets[rank_idx] + feature_idx * batch_size_per_rank[rank_idx] + b52 for feature_idx in range(local_split)53 for rank_idx in feature_order54 for b in range(batch_size_per_rank[rank_idx])55 ]56 return torch.tensor(recat, device=device, dtype=torch.int32)57 58 59def _permute_segments(60 data: torch.Tensor,61 segment_lengths: torch.Tensor,62 recat: torch.Tensor,63) -> torch.Tensor:64 segment_lengths = segment_lengths.to(device=data.device, dtype=torch.long)65 offsets = torch.zeros(66 segment_lengths.numel() + 1, dtype=torch.long, device=data.device67 )68 offsets[1:] = torch.cumsum(segment_lengths, dim=0)69 chunks = [70 data[int(offsets[idx].item()) : int(offsets[idx + 1].item())]71 for idx in recat.long().tolist()72 ]73 return torch.cat(chunks, dim=0) if chunks else data.new_empty((0,))74 75 76def _permute_2d_sparse_data(77 recat: torch.Tensor,78 lengths_2d: torch.Tensor,79 values: torch.Tensor,80 weights: Optional[torch.Tensor],81) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:82 recat = recat.long()83 row_lengths = lengths_2d.sum(dim=1).to(torch.long)84 lengths_out = lengths_2d[recat]85 values_out = _permute_segments(values, row_lengths, recat)86 weights_out = None87 if weights is not None:88 weights_out = _permute_segments(weights, row_lengths, recat)89 return lengths_out, values_out, weights_out90 91 92def _all_to_all_tensor(93 tensor: torch.Tensor,94 input_splits: List[int],95 output_splits: List[int],96 pg: dist.ProcessGroup,97) -> Tuple[torch.Tensor, dist.Work]:98 output = torch.empty(99 (sum(output_splits),), dtype=tensor.dtype, device=tensor.device100 )101 work = dist.all_to_all_single(102 output,103 tensor,104 output_split_sizes=output_splits,105 input_split_sizes=input_splits,106 group=pg,107 async_op=True,108 )109 return output, work110 111 112@torch.no_grad()113def solution(114 lengths: torch.Tensor,115 values: torch.Tensor,116 key_splits: List[int],117 batch_size: int,118 pg: Optional[dist.ProcessGroup] = None,119 weights: Optional[torch.Tensor] = None,120 stride_per_key: Optional[List[int]] = None,121 stagger: int = 1,122) -> Dict[str, torch.Tensor]:123 pg = pg or dist.group.WORLD124 world_size = dist.get_world_size(pg)125 rank = dist.get_rank(pg)126 device = lengths.device127 128 num_features = sum(key_splits)129 variable_stride = stride_per_key is not None130 if stride_per_key is None:131 stride_per_key = [batch_size] * num_features132 133 length_per_key = _lengths_per_key(lengths, stride_per_key)134 length_splits = _sum_by_splits(stride_per_key, key_splits)135 value_splits = _sum_by_splits(length_per_key, key_splits)136 137 input_splits = [length_splits, value_splits]138 input_tensors = [lengths, values]139 if variable_stride:140 input_splits.append(key_splits)141 input_tensors.append(142 torch.tensor(stride_per_key, dtype=torch.long, device=device)143 )144 if weights is not None:145 input_splits.append(value_splits)146 input_tensors.append(weights)147 148 split_tensors = [149 torch.tensor(splits, dtype=torch.long, device=device) for splits in input_splits150 ]151 if not variable_stride:152 split_tensors.append(153 torch.full((world_size,), batch_size, dtype=torch.long, device=device)154 )155 156 meta_input = torch.stack(split_tensors, dim=1).flatten()157 meta_output = torch.empty_like(meta_input)158 dist.all_to_all_single(meta_output, meta_input, group=pg)159 meta_rows = [160 [int(item) for item in row]161 for row in meta_output.view(world_size, -1).T.tolist()162 ]163 if variable_stride:164 output_splits = meta_rows165 stride_per_rank = None166 else:167 output_splits = meta_rows[:-1]168 stride_per_rank = meta_rows[-1]169 170 outputs: List[torch.Tensor] = []171 works: List[dist.Work] = []172 for tensor, in_splits, out_splits in zip(173 input_tensors, input_splits, output_splits174 ):175 output, work = _all_to_all_tensor(tensor, in_splits, out_splits, pg)176 outputs.append(output)177 works.append(work)178 for work in works:179 work.wait()180 181 recv_lengths = outputs[0]182 recv_values = outputs[1]183 recv_strides: Optional[torch.Tensor] = outputs[2] if variable_stride else None184 recv_weights: Optional[torch.Tensor] = None185 if weights is not None:186 recv_weights = outputs[-1]187 188 local_split = key_splits[rank]189 if variable_stride:190 assert recv_strides is not None191 recat = _get_recat(local_split, world_size, stagger, device=device)192 if recat is not None:193 value_segment_lengths = torch.tensor(194 _lengths_per_key(recv_lengths, recv_strides.to(torch.long).tolist()),195 dtype=torch.long,196 device=device,197 )198 recv_lengths = _permute_segments(recv_lengths, recv_strides, recat)199 recv_values = _permute_segments(recv_values, value_segment_lengths, recat)200 if recv_weights is not None:201 recv_weights = _permute_segments(202 recv_weights, value_segment_lengths, recat203 )204 stride_per_key_per_rank = recv_strides.view(world_size, local_split).T205 if stagger > 1:206 order = (207 torch.arange(world_size, device=device)208 .view(stagger, -1)209 .T.reshape(-1)210 )211 stride_per_key_per_rank = stride_per_key_per_rank[:, order]212 result: Dict[str, torch.Tensor] = {213 "lengths": recv_lengths,214 "values": recv_values,215 "stride_per_key_per_rank": stride_per_key_per_rank,216 }217 else:218 assert stride_per_rank is not None219 single_batch_per_rank = all(220 stride == stride_per_rank[0] for stride in stride_per_rank221 )222 if single_batch_per_rank:223 recat = _get_recat(local_split, world_size, stagger, device=device)224 if recat is not None and stride_per_rank[0] > 0:225 lengths_2d, recv_values, recv_weights = _permute_2d_sparse_data(226 recat,227 recv_lengths.view(-1, stride_per_rank[0]),228 recv_values,229 recv_weights,230 )231 recv_lengths = lengths_2d.reshape(-1)232 else:233 recat = _get_recat(234 local_split,235 world_size,236 stagger,237 device=device,238 batch_size_per_rank=stride_per_rank,239 )240 if recat is not None:241 recv_values = _permute_segments(recv_values, recv_lengths, recat)242 if recv_weights is not None:243 recv_weights = _permute_segments(recv_weights, recv_lengths, recat)244 recv_lengths = recv_lengths[recat.long()]245 result = {246 "lengths": recv_lengths,247 "values": recv_values,248 "stride": torch.tensor(sum(stride_per_rank), device=device),249 "stride_per_rank": torch.tensor(stride_per_rank, device=device),250 }251 252 if recv_weights is not None:253 result["weights"] = recv_weights254 return result