togethercomputer/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. Inputs are deterministic — reproduce them with create_input_tensor(rank, world_size, problem_id, base_shape, dtype, trial) from that file; you do not need stored .pt files. Files Path Description… See the full description on the dataset page: https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems.
0177
1from typing import List, Optional, Tuple2 3import numpy as np4import torch5import torch.distributed as dist6 7 8def _sample_one_hop_csc_dist(9 input_nodes: torch.Tensor,10 k: int,11 colptr: torch.Tensor,12 row: torch.Tensor,13 replace: bool = False,14) -> Tuple[torch.Tensor, torch.Tensor, List[int]]:15 n = input_nodes.numel()16 sampled_nodes = []17 sampled_edges = []18 cumsum = [n]19 20 for i in range(n):21 v = int(input_nodes[i].item())22 start = int(colptr[v].item())23 end = int(colptr[v + 1].item())24 deg = end - start25 take = min(k, deg) if k >= 0 else deg26 27 if take > 0:28 perm = torch.arange(take, device=input_nodes.device)29 sampled_nodes.append(row[start:end].index_select(0, perm))30 sampled_edges.append(torch.arange(start, end, device=input_nodes.device).index_select(0, perm))31 32 cumsum.append(cumsum[-1] + take)33 34 nbr_tensor = (35 torch.cat(sampled_nodes)36 if sampled_nodes37 else torch.empty(0, dtype=torch.long, device=input_nodes.device)38 )39 eid_tensor = (40 torch.cat(sampled_edges)41 if sampled_edges42 else torch.empty(0, dtype=torch.long, device=input_nodes.device)43 )44 return torch.cat([input_nodes, nbr_tensor]), eid_tensor, cumsum45 46 47def _remove_duplicates(out_node: torch.Tensor, node: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:48 num_nodes = node.numel()49 node_combined = torch.cat([node, out_node])50 _, idx = np.unique(node_combined.cpu().numpy(), return_index=True)51 idx = torch.from_numpy(idx).to(node.device).sort().values52 node = node_combined[idx]53 src = node[num_nodes:]54 return src, node55 56 57def _relabel_neighborhood(58 node: torch.Tensor,59 dst_with_dupl: torch.Tensor,60 node_with_dupl: torch.Tensor,61) -> Tuple[torch.Tensor, torch.Tensor]:62 if node_with_dupl.numel() == 0:63 return node.new_empty(0), node.new_empty(0)64 65 assoc = torch.full(66 (int(node.max().item()) + 1,),67 -1,68 dtype=torch.long,69 device=node.device,70 )71 assoc[node] = torch.arange(node.numel(), device=node.device)72 row = assoc[node_with_dupl]73 col = assoc[dst_with_dupl]74 return row, col75 76 77def _exchange_nodes(78 send_nodes_list: List[torch.Tensor],79 group: dist.ProcessGroup,80) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:81 world_size = dist.get_world_size(group)82 device = send_nodes_list[0].device83 send_counts = torch.tensor([x.numel() for x in send_nodes_list], dtype=torch.long, device=device)84 recv_counts = torch.empty_like(send_counts)85 dist.all_to_all_single(recv_counts, send_counts, group=group)86 87 send_nodes = torch.cat(send_nodes_list) if send_nodes_list else torch.empty(0, dtype=torch.long, device=device)88 recv_nodes = torch.empty(int(recv_counts.sum().item()), dtype=torch.long, device=device)89 dist.all_to_all_single(90 recv_nodes,91 send_nodes,92 input_split_sizes=send_counts.cpu().tolist(),93 output_split_sizes=recv_counts.cpu().tolist(),94 group=group,95 )96 return recv_nodes, send_counts, recv_counts97 98 99def _exchange_replies(100 sampled_nodes: torch.Tensor,101 sampled_edges: torch.Tensor,102 sampled_counts: torch.Tensor,103 recv_counts: torch.Tensor,104 group: dist.ProcessGroup,105) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:106 world_size = dist.get_world_size(group)107 device = sampled_nodes.device108 recv_splits = recv_counts.cpu().tolist()109 send_node_counts = torch.empty(world_size, dtype=torch.long, device=device)110 offset = 0111 for r, count in enumerate(recv_splits):112 send_node_counts[r] = sampled_counts[offset : offset + count].sum()113 offset += count114 115 reply_node_counts = torch.empty_like(send_node_counts)116 dist.all_to_all_single(reply_node_counts, send_node_counts, group=group)117 118 reply_count_counts = torch.empty_like(recv_counts)119 dist.all_to_all_single(reply_count_counts, recv_counts, group=group)120 121 reply_nodes = torch.empty(int(reply_node_counts.sum().item()), dtype=torch.long, device=device)122 reply_edges = torch.empty_like(reply_nodes)123 reply_counts = torch.empty(int(reply_count_counts.sum().item()), dtype=torch.long, device=device)124 125 dist.all_to_all_single(126 reply_nodes,127 sampled_nodes,128 input_split_sizes=send_node_counts.cpu().tolist(),129 output_split_sizes=reply_node_counts.cpu().tolist(),130 group=group,131 )132 dist.all_to_all_single(133 reply_edges,134 sampled_edges,135 input_split_sizes=send_node_counts.cpu().tolist(),136 output_split_sizes=reply_node_counts.cpu().tolist(),137 group=group,138 )139 dist.all_to_all_single(140 reply_counts,141 sampled_counts,142 input_split_sizes=recv_splits,143 output_split_sizes=reply_count_counts.cpu().tolist(),144 group=group,145 )146 return reply_nodes, reply_edges, reply_counts147 148 149@torch.no_grad()150def solution(151 seed_nodes: torch.Tensor,152 fanouts: List[int],153 local_adj_row_ptr: torch.Tensor,154 local_adj_col: torch.Tensor,155 node_to_rank: torch.Tensor,156 group: Optional[dist.ProcessGroup] = None,157 replace: bool = False,158) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:159 group = group or dist.group.WORLD160 world_size = dist.get_world_size(group)161 device = seed_nodes.device162 163 seed = seed_nodes.to(dtype=torch.long, device=device)164 src = seed.clone()165 node = src.clone()166 node_with_dupl = [seed.new_empty(0)]167 dst_with_dupl = [seed.new_empty(0)]168 edge = [seed.new_empty(0)]169 170 for fanout in fanouts:171 if src.numel() == 0:172 break173 174 partition_ids = node_to_rank[src].to(torch.long)175 partition_orders = torch.empty_like(partition_ids)176 send_nodes_list = []177 send_pos_list = []178 for r in range(world_size):179 pos = (partition_ids == r).nonzero(as_tuple=False).flatten()180 partition_orders[pos] = torch.arange(pos.numel(), dtype=torch.long, device=device)181 send_nodes_list.append(src[pos])182 send_pos_list.append(pos)183 184 recv_nodes, send_counts, recv_counts = _exchange_nodes(send_nodes_list, group)185 node_out, edge_out, cumsum = _sample_one_hop_csc_dist(186 recv_nodes, int(fanout), local_adj_row_ptr, local_adj_col, replace187 )188 189 seed_size = recv_nodes.numel()190 sampled_nodes = node_out[seed_size:]191 sampled_counts = torch.tensor(192 np.subtract(np.array(cumsum[1:]), np.array(cumsum[:-1])),193 dtype=torch.long,194 device=device,195 )196 197 reply_nodes, reply_edges, reply_counts = _exchange_replies(198 sampled_nodes, edge_out, sampled_counts, recv_counts, group199 )200 201 rank_offsets = torch.cat(202 [send_counts.new_zeros(1), torch.cumsum(send_counts, dim=0)[:-1]]203 )204 grouped_index = rank_offsets[partition_ids] + partition_orders205 node_chunks = list(torch.split(reply_nodes, reply_counts.cpu().tolist()))206 edge_chunks = list(torch.split(reply_edges, reply_counts.cpu().tolist()))207 208 ordered_nodes = []209 ordered_edges = []210 ordered_dst = []211 for idx in grouped_index.tolist():212 ordered_nodes.append(node_chunks[idx])213 ordered_edges.append(edge_chunks[idx])214 for dst_node, count in zip(src, reply_counts[grouped_index]):215 ordered_dst.append(dst_node.repeat(int(count.item())))216 217 out_node = torch.cat(ordered_nodes) if ordered_nodes else seed.new_empty(0)218 out_edge = torch.cat(ordered_edges) if ordered_edges else seed.new_empty(0)219 out_dst = torch.cat(ordered_dst) if ordered_dst else seed.new_empty(0)220 if out_node.numel() == 0:221 break222 223 src, node = _remove_duplicates(out_node, node)224 node_with_dupl.append(out_node)225 dst_with_dupl.append(out_dst)226 edge.append(out_edge)227 228 node_dupl = torch.cat(node_with_dupl)229 dst_dupl = torch.cat(dst_with_dupl)230 row, col = _relabel_neighborhood(node, dst_dupl, node_dupl)231 return node, row, col, torch.cat(edge)