Team Ai
Datasetpublic

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.

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes177downloads
64_gnn_neighbor_sampling.py231 linesDownload Raw Back to reference
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)