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
49_moe_ep_balanced.py317 linesDownload Raw Back to reference
1# Expert-parallel (EP) fused MoE forward — BALANCED EP.2#3# Same kernel as problem 31 (router -> permute -> all_to_all dispatch -> per-expert4# SiLU MLP -> all_to_all combine -> unpermute), but the harness sets5# num_experts == world_size, i.e. exactly one expert per rank. With uniform routing6# this gives the balanced all_to_all dispatch pattern (each rank sends/receives a7# roughly equal token count).8from typing import List, Optional, Tuple, Union9 10import torch11import torch.distributed as dist12 13 14class _AllToAll(torch.autograd.Function):15    @staticmethod16    def forward(ctx, group, input, output_split_sizes, input_split_sizes):17        ctx.group = group18        ctx.output_split_sizes = output_split_sizes19        ctx.input_split_sizes = input_split_sizes20        if dist.get_world_size(group=group) == 1:21            return input.contiguous()22        input = input.contiguous()23        if output_split_sizes is None:24            output = torch.empty_like(input)25        else:26            output = torch.empty(27                size=(sum(output_split_sizes), input.size(1)),28                dtype=input.dtype,29                device=input.device,30            )31        dist.all_to_all_single(32            output,33            input,34            output_split_sizes=output_split_sizes,35            input_split_sizes=input_split_sizes,36            group=group,37        )38        return output39 40    @staticmethod41    def backward(ctx, grad_output):42        return (43            None,44            _AllToAll.apply(45                ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes46            ),47            None,48            None,49        )50 51 52def _all_to_all(53    group: dist.ProcessGroup,54    input: torch.Tensor,55    output_split_sizes: Optional[List[int]],56    input_split_sizes: Optional[List[int]],57) -> torch.Tensor:58    return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)59 60 61def _preprocess(62    expert_mask: torch.Tensor,63    num_experts: int,64    ep_group: dist.ProcessGroup,65) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:66    ep_size = ep_group.size()67    num_local_experts = num_experts // ep_size68    rank = dist.get_rank(ep_group)69    num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))70    input_splits = (71        num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()72    )73    num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)74    output_size = ep_size * num_local_tokens_per_expert_flat.numel()75    num_global_tokens_per_expert_flat = torch.empty(76        output_size,77        dtype=num_local_tokens_per_expert.dtype,78        device=num_local_tokens_per_expert.device,79    )80    dist.all_gather_into_tensor(81        num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group82    )83    num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(84        ep_size, num_local_tokens_per_expert.size(0)85    )86    start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts87    num_global_tokens_per_local_expert = num_global_tokens_per_expert[88        :, start_idx:end_idx89    ].contiguous()90    output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()91    num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(92        dim=093    ).to(torch.device("cpu"), non_blocking=True)94    num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(95        -1, num_local_experts96    ).to(torch.device("cpu"), non_blocking=True)97    return (98        input_splits,99        output_splits,100        num_global_tokens_per_local_expert,101        num_global_sum_tokens_per_local_expert,102    )103 104 105def _permute(106    tokens: torch.Tensor, routing_map: torch.Tensor107) -> Tuple[torch.Tensor, torch.Tensor]:108    num_tokens, _ = tokens.shape109    num_experts = routing_map.shape[0]110    routing_map = routing_map.bool()111    token_indices = (112        torch.arange(num_tokens, device=routing_map.device)113        .unsqueeze(0)114        .expand(num_experts, -1)115    )116    sorted_indices = token_indices.masked_select(routing_map)117    permuted_input = tokens.index_select(0, sorted_indices)118    return permuted_input, sorted_indices119 120 121def _sort_chunks_by_idxs(122    input: torch.Tensor,123    split_sizes: Union[torch.Tensor, List[int]],124    sorted_idxs: List[int],125) -> torch.Tensor:126    if isinstance(split_sizes, torch.Tensor):127        split_sizes = split_sizes.tolist()128    chunks = torch.split(input, split_sizes, dim=0)129    return torch.cat([chunks[i] for i in sorted_idxs], dim=0)130 131 132def _generate_weights_idx(133    routing_weights: torch.Tensor,134    selected_experts: torch.Tensor,135    num_experts: int,136) -> torch.Tensor:137    num_tokens, topk = routing_weights.shape138    weights_idx = torch.zeros(139        (num_tokens, num_experts),140        dtype=routing_weights.dtype,141        device=routing_weights.device,142    )143    weights_idx.scatter_add_(1, selected_experts, routing_weights)144    return weights_idx145 146 147def _unpermute(148    tokens: torch.Tensor,149    routing_weights: torch.Tensor,150    hidden_states_shape: torch.Size,151    permutation_mapping: torch.Tensor,152    routing_map: torch.Tensor,153) -> torch.Tensor:154    tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())155    tokens = tokens * tokens_weight.unsqueeze(-1)156    hidden_dim = hidden_states_shape[-1]157    unpermuted_tokens = torch.zeros(158        hidden_states_shape, device=tokens.device, dtype=tokens.dtype159    )160    expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)161    unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)162    return unpermuted_tokens163 164 165def token_pre_all2all(166    hidden_states: torch.Tensor,167    expert_mask: torch.Tensor,168    num_experts: int,169    input_splits: List[int],170    output_splits: List[int],171    num_global_tokens_per_local_expert: torch.Tensor,172    group: Optional[dist.ProcessGroup] = None,173) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:174    group = group or dist.group.WORLD175    hidden_dim = hidden_states.size(-1)176    hidden_states = hidden_states.reshape(-1, hidden_dim)177    org_hidden_states_shape = hidden_states.shape178    routing_map = expert_mask.sum(dim=1)179 180    local_permuted_hidden_states, local_input_permutation_mapping = _permute(181        hidden_states, routing_map182    )183    expected_tokens = sum(input_splits)184    actual_tokens = local_permuted_hidden_states.shape[0]185    if expected_tokens != actual_tokens:186        raise RuntimeError(187            f"EP split mismatch: input_splits sum ({expected_tokens}) != "188            f"permuted tokens ({actual_tokens})"189        )190 191    global_permuted_hidden_states = _all_to_all(192        group, local_permuted_hidden_states, output_splits, input_splits193    )194    num_local_experts = num_experts // dist.get_world_size(group)195    permute_order = (196        torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()197    )198    split_sizes = num_global_tokens_per_local_expert.ravel().tolist()199    global_permuted_hidden_states = _sort_chunks_by_idxs(200        global_permuted_hidden_states, split_sizes, permute_order201    )202    return (203        global_permuted_hidden_states,204        routing_map,205        local_input_permutation_mapping,206        org_hidden_states_shape,207    )208 209 210def tokens_post_all2all(211    expert_outputs: torch.Tensor,212    routing_weights: torch.Tensor,213    selected_experts: torch.Tensor,214    num_experts: int,215    input_splits: List[int],216    output_splits: List[int],217    num_global_tokens_per_local_expert: torch.Tensor,218    routing_map: torch.Tensor,219    local_input_permutation_mapping: torch.Tensor,220    org_hidden_states_shape: torch.Size,221    group: Optional[dist.ProcessGroup] = None,222) -> torch.Tensor:223    group = group or dist.group.WORLD224    num_local_experts = num_experts // dist.get_world_size(group)225    unpermute_order = (226        torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()227    )228    split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()229    expert_outputs = _sort_chunks_by_idxs(230        expert_outputs, split_sizes, unpermute_order231    )232    unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)233    weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)234    unpermute_outputs = _unpermute(235        unpermute_outputs,236        weights_idx,237        org_hidden_states_shape,238        local_input_permutation_mapping,239        routing_map,240    )241    return unpermute_outputs242 243 244def expert_forward(245    x: torch.Tensor,246    gate_proj: torch.nn.Linear,247    up_proj: torch.nn.Linear,248    down_proj: torch.nn.Linear,249) -> torch.Tensor:250    gate = torch.nn.functional.silu(gate_proj(x))251    up = up_proj(x)252    return down_proj(gate * up)253 254 255def solution(256    hidden_states: torch.Tensor,257    gate_weight: torch.Tensor,258    gate_bias: Optional[torch.Tensor],259    gate_proj: torch.nn.Linear,260    up_proj: torch.nn.Linear,261    down_proj: torch.nn.Linear,262    num_experts: int,263    top_k: int,264    group: Optional[dist.ProcessGroup] = None,265) -> torch.Tensor:266    group = group or dist.group.WORLD267    hidden_dim = hidden_states.size(-1)268    num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)269 270    router_logits = torch.nn.functional.linear(271        hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias272    )273    routing_weights, selected_experts = torch.topk(274        torch.softmax(router_logits, dim=-1), top_k, dim=-1275    )276    expert_mask = torch.nn.functional.one_hot(277        selected_experts, num_classes=num_experts278    ).permute(2, 1, 0)279 280    input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(281        expert_mask, num_experts, group282    )283 284    (285        global_permuted_hidden_states,286        routing_map,287        local_input_permutation_mapping,288        org_hidden_states_shape,289    ) = token_pre_all2all(290        hidden_states,291        expert_mask,292        num_experts,293        input_splits,294        output_splits,295        num_global_tokens_per_local_expert,296        group,297    )298 299    expert_outputs = expert_forward(300        global_permuted_hidden_states, gate_proj, up_proj, down_proj301    )302 303    out = tokens_post_all2all(304        expert_outputs,305        routing_weights,306        selected_experts,307        num_experts,308        input_splits,309        output_splits,310        num_global_tokens_per_local_expert,311        routing_map,312        local_input_permutation_mapping,313        org_hidden_states_shape,314        group,315    )316    return out317