CoolFace
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 4mo agoView on Hugging Face
0likes263downloads
70_torchrec_kjt_all2all.py254 linesDownload Raw Back to reference
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