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
31_fused_moe_fwd.py316 linesDownload Raw Back to reference
1# Expert-parallel (EP) fused MoE forward — BASE case.2#3# This is the generic fused MoE forward pass: router (softmax + top-k) -> token4# permutation -> all_to_all dispatch -> per-expert SiLU MLP -> all_to_all combine ->5# weighted unpermute. Here the expert count is fixed (num_experts = 8) regardless of6# world size, so the EP load pattern depends on the launch configuration.7from typing import List, Optional, Tuple, Union8 9import torch10import torch.distributed as dist11 12 13class _AllToAll(torch.autograd.Function):14    @staticmethod15    def forward(ctx, group, input, output_split_sizes, input_split_sizes):16        ctx.group = group17        ctx.output_split_sizes = output_split_sizes18        ctx.input_split_sizes = input_split_sizes19        if dist.get_world_size(group=group) == 1:20            return input.contiguous()21        input = input.contiguous()22        if output_split_sizes is None:23            output = torch.empty_like(input)24        else:25            output = torch.empty(26                size=(sum(output_split_sizes), input.size(1)),27                dtype=input.dtype,28                device=input.device,29            )30        dist.all_to_all_single(31            output,32            input,33            output_split_sizes=output_split_sizes,34            input_split_sizes=input_split_sizes,35            group=group,36        )37        return output38 39    @staticmethod40    def backward(ctx, grad_output):41        return (42            None,43            _AllToAll.apply(44                ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes45            ),46            None,47            None,48        )49 50 51def _all_to_all(52    group: dist.ProcessGroup,53    input: torch.Tensor,54    output_split_sizes: Optional[List[int]],55    input_split_sizes: Optional[List[int]],56) -> torch.Tensor:57    return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)58 59 60def _preprocess(61    expert_mask: torch.Tensor,62    num_experts: int,63    ep_group: dist.ProcessGroup,64) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:65    ep_size = ep_group.size()66    num_local_experts = num_experts // ep_size67    rank = dist.get_rank(ep_group)68    num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))69    input_splits = (70        num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()71    )72    num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)73    output_size = ep_size * num_local_tokens_per_expert_flat.numel()74    num_global_tokens_per_expert_flat = torch.empty(75        output_size,76        dtype=num_local_tokens_per_expert.dtype,77        device=num_local_tokens_per_expert.device,78    )79    dist.all_gather_into_tensor(80        num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group81    )82    num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(83        ep_size, num_local_tokens_per_expert.size(0)84    )85    start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts86    num_global_tokens_per_local_expert = num_global_tokens_per_expert[87        :, start_idx:end_idx88    ].contiguous()89    output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()90    num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(91        dim=092    ).to(torch.device("cpu"), non_blocking=True)93    num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(94        -1, num_local_experts95    ).to(torch.device("cpu"), non_blocking=True)96    return (97        input_splits,98        output_splits,99        num_global_tokens_per_local_expert,100        num_global_sum_tokens_per_local_expert,101    )102 103 104def _permute(105    tokens: torch.Tensor, routing_map: torch.Tensor106) -> Tuple[torch.Tensor, torch.Tensor]:107    num_tokens, _ = tokens.shape108    num_experts = routing_map.shape[0]109    routing_map = routing_map.bool()110    token_indices = (111        torch.arange(num_tokens, device=routing_map.device)112        .unsqueeze(0)113        .expand(num_experts, -1)114    )115    sorted_indices = token_indices.masked_select(routing_map)116    permuted_input = tokens.index_select(0, sorted_indices)117    return permuted_input, sorted_indices118 119 120def _sort_chunks_by_idxs(121    input: torch.Tensor,122    split_sizes: Union[torch.Tensor, List[int]],123    sorted_idxs: List[int],124) -> torch.Tensor:125    if isinstance(split_sizes, torch.Tensor):126        split_sizes = split_sizes.tolist()127    chunks = torch.split(input, split_sizes, dim=0)128    return torch.cat([chunks[i] for i in sorted_idxs], dim=0)129 130 131def _generate_weights_idx(132    routing_weights: torch.Tensor,133    selected_experts: torch.Tensor,134    num_experts: int,135) -> torch.Tensor:136    num_tokens, topk = routing_weights.shape137    weights_idx = torch.zeros(138        (num_tokens, num_experts),139        dtype=routing_weights.dtype,140        device=routing_weights.device,141    )142    weights_idx.scatter_add_(1, selected_experts, routing_weights)143    return weights_idx144 145 146def _unpermute(147    tokens: torch.Tensor,148    routing_weights: torch.Tensor,149    hidden_states_shape: torch.Size,150    permutation_mapping: torch.Tensor,151    routing_map: torch.Tensor,152) -> torch.Tensor:153    tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())154    tokens = tokens * tokens_weight.unsqueeze(-1)155    hidden_dim = hidden_states_shape[-1]156    unpermuted_tokens = torch.zeros(157        hidden_states_shape, device=tokens.device, dtype=tokens.dtype158    )159    expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)160    unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)161    return unpermuted_tokens162 163 164def token_pre_all2all(165    hidden_states: torch.Tensor,166    expert_mask: torch.Tensor,167    num_experts: int,168    input_splits: List[int],169    output_splits: List[int],170    num_global_tokens_per_local_expert: torch.Tensor,171    group: Optional[dist.ProcessGroup] = None,172) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:173    group = group or dist.group.WORLD174    hidden_dim = hidden_states.size(-1)175    hidden_states = hidden_states.reshape(-1, hidden_dim)176    org_hidden_states_shape = hidden_states.shape177    routing_map = expert_mask.sum(dim=1)178 179    local_permuted_hidden_states, local_input_permutation_mapping = _permute(180        hidden_states, routing_map181    )182    expected_tokens = sum(input_splits)183    actual_tokens = local_permuted_hidden_states.shape[0]184    if expected_tokens != actual_tokens:185        raise RuntimeError(186            f"EP split mismatch: input_splits sum ({expected_tokens}) != "187            f"permuted tokens ({actual_tokens})"188        )189 190    global_permuted_hidden_states = _all_to_all(191        group, local_permuted_hidden_states, output_splits, input_splits192    )193    num_local_experts = num_experts // dist.get_world_size(group)194    permute_order = (195        torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()196    )197    split_sizes = num_global_tokens_per_local_expert.ravel().tolist()198    global_permuted_hidden_states = _sort_chunks_by_idxs(199        global_permuted_hidden_states, split_sizes, permute_order200    )201    return (202        global_permuted_hidden_states,203        routing_map,204        local_input_permutation_mapping,205        org_hidden_states_shape,206    )207 208 209def tokens_post_all2all(210    expert_outputs: torch.Tensor,211    routing_weights: torch.Tensor,212    selected_experts: torch.Tensor,213    num_experts: int,214    input_splits: List[int],215    output_splits: List[int],216    num_global_tokens_per_local_expert: torch.Tensor,217    routing_map: torch.Tensor,218    local_input_permutation_mapping: torch.Tensor,219    org_hidden_states_shape: torch.Size,220    group: Optional[dist.ProcessGroup] = None,221) -> torch.Tensor:222    group = group or dist.group.WORLD223    num_local_experts = num_experts // dist.get_world_size(group)224    unpermute_order = (225        torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()226    )227    split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()228    expert_outputs = _sort_chunks_by_idxs(229        expert_outputs, split_sizes, unpermute_order230    )231    unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)232    weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)233    unpermute_outputs = _unpermute(234        unpermute_outputs,235        weights_idx,236        org_hidden_states_shape,237        local_input_permutation_mapping,238        routing_map,239    )240    return unpermute_outputs241 242 243def expert_forward(244    x: torch.Tensor,245    gate_proj: torch.nn.Linear,246    up_proj: torch.nn.Linear,247    down_proj: torch.nn.Linear,248) -> torch.Tensor:249    gate = torch.nn.functional.silu(gate_proj(x))250    up = up_proj(x)251    return down_proj(gate * up)252 253 254def solution(255    hidden_states: torch.Tensor,256    gate_weight: torch.Tensor,257    gate_bias: Optional[torch.Tensor],258    gate_proj: torch.nn.Linear,259    up_proj: torch.nn.Linear,260    down_proj: torch.nn.Linear,261    num_experts: int,262    top_k: int,263    group: Optional[dist.ProcessGroup] = None,264) -> torch.Tensor:265    group = group or dist.group.WORLD266    hidden_dim = hidden_states.size(-1)267    num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)268 269    router_logits = torch.nn.functional.linear(270        hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias271    )272    routing_weights, selected_experts = torch.topk(273        torch.softmax(router_logits, dim=-1), top_k, dim=-1274    )275    expert_mask = torch.nn.functional.one_hot(276        selected_experts, num_classes=num_experts277    ).permute(2, 1, 0)278 279    input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(280        expert_mask, num_experts, group281    )282 283    (284        global_permuted_hidden_states,285        routing_map,286        local_input_permutation_mapping,287        org_hidden_states_shape,288    ) = token_pre_all2all(289        hidden_states,290        expert_mask,291        num_experts,292        input_splits,293        output_splits,294        num_global_tokens_per_local_expert,295        group,296    )297 298    expert_outputs = expert_forward(299        global_permuted_hidden_states, gate_proj, up_proj, down_proj300    )301 302    out = tokens_post_all2all(303        expert_outputs,304        routing_weights,305        selected_experts,306        num_experts,307        input_splits,308        output_splits,309        num_global_tokens_per_local_expert,310        routing_map,311        local_input_permutation_mapping,312        org_hidden_states_shape,313        group,314    )315    return out316