Team Ai
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 5mo agoView on Hugging Face
0likes123downloads
29_moe_post_all2all.py102 linesDownload Raw Back to reference
1from typing import List, Optional, Union2 3import torch4import torch.distributed as dist5 6 7def _sort_chunks_by_idxs(8    input: torch.Tensor,9    split_sizes: Union[torch.Tensor, List[int]],10    sorted_idxs: List[int],11) -> torch.Tensor:12    if isinstance(split_sizes, torch.Tensor):13        split_sizes = split_sizes.tolist()14    chunks = torch.split(input, split_sizes, dim=0)15    return torch.cat([chunks[i] for i in sorted_idxs], dim=0)16 17 18def _all_to_all_forward(19    group: dist.ProcessGroup,20    input: torch.Tensor,21    output_split_sizes: Optional[List[int]],22    input_split_sizes: Optional[List[int]],23) -> torch.Tensor:24    if dist.get_world_size(group) == 1:25        return input.contiguous()26    input = input.contiguous()27    out_size = sum(output_split_sizes) if output_split_sizes else input.size(0)28    output = torch.empty((out_size, input.size(1)), dtype=input.dtype, device=input.device)29    dist.all_to_all_single(30        output, input,31        output_split_sizes=output_split_sizes,32        input_split_sizes=input_split_sizes,33        group=group,34    )35    return output36 37 38def _generate_weights_idx(39    routing_weights: torch.Tensor,40    selected_experts: torch.Tensor,41    num_experts: int,42) -> torch.Tensor:43    num_tokens, topk = routing_weights.shape44    weights_idx = torch.zeros(45        (num_tokens, num_experts), dtype=routing_weights.dtype, device=routing_weights.device46    )47    weights_idx.scatter_add_(1, selected_experts, routing_weights)48    return weights_idx49 50 51def _unpermute(52    tokens: torch.Tensor,53    routing_weights: torch.Tensor,54    hidden_states_shape: torch.Size,55    permutation_mapping: torch.Tensor,56    routing_map: torch.Tensor,57) -> torch.Tensor:58    tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())59    tokens = tokens * tokens_weight.unsqueeze(-1)60    hidden_dim = hidden_states_shape[-1]61    unpermuted_tokens = torch.zeros(hidden_states_shape, device=tokens.device, dtype=tokens.dtype)62    expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)63    unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)64    return unpermuted_tokens65 66 67def solution(68    expert_outputs: torch.Tensor,69    routing_weights: torch.Tensor,70    selected_experts: torch.Tensor,71    num_experts: int,72    input_splits: Union[List[int], torch.Tensor],73    output_splits: Union[List[int], torch.Tensor],74    num_global_tokens_per_local_expert: torch.Tensor,75    routing_map: torch.Tensor,76    local_input_permutation_mapping: torch.Tensor,77    org_hidden_states_shape: torch.Size,78    group: Optional[dist.ProcessGroup] = None,79) -> torch.Tensor:80    group = group or dist.group.WORLD81    num_local_experts = num_experts // dist.get_world_size(group)82    unpermute_order = torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()83 84    expert_outputs = _sort_chunks_by_idxs(85        expert_outputs,86        num_global_tokens_per_local_expert.T.ravel(),87        unpermute_order,88    )89 90    unpermute_outputs = _all_to_all_forward(group, expert_outputs, input_splits, output_splits)91 92    weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)93    unpermute_outputs = _unpermute(94        unpermute_outputs,95        weights_idx,96        org_hidden_states_shape,97        local_input_permutation_mapping,98        routing_map,99    )100 101    return unpermute_outputs102