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
51_moe_ep_narrow.py338 linesDownload Raw Back to reference
1from typing import List, Optional, Tuple, Union2 3import torch4import torch.distributed as dist5 6_EP_SUBGROUP_CACHE: dict[tuple[int, int], None | list] = {}7 8 9def _resolve_ep_group_for_narrow_moe(num_experts: int) -> dist.ProcessGroup:10    if not dist.is_initialized():11        raise RuntimeError("torch.distributed must be initialized")12    ws = dist.get_world_size()13    rank = dist.get_rank()14    key = (ws, num_experts)15    if key not in _EP_SUBGROUP_CACHE:16        if num_experts >= ws:17            _EP_SUBGROUP_CACHE[key] = None18        elif ws % num_experts != 0:19            raise ValueError(20                f"narrow EP requires world_size ({ws}) % num_experts ({num_experts}) == 0"21            )22        else:23            groups: list = []24            for r in range(ws // num_experts):25                ranks = list(range(r * num_experts, (r + 1) * num_experts))26                groups.append(dist.new_group(ranks))27            _EP_SUBGROUP_CACHE[key] = groups28    entry = _EP_SUBGROUP_CACHE[key]29    if entry is None:30        return dist.group.WORLD31    return entry[rank // num_experts]32 33 34class _AllToAll(torch.autograd.Function):35    @staticmethod36    def forward(ctx, group, input, output_split_sizes, input_split_sizes):37        ctx.group = group38        ctx.output_split_sizes = output_split_sizes39        ctx.input_split_sizes = input_split_sizes40        if dist.get_world_size(group=group) == 1:41            return input.contiguous()42        input = input.contiguous()43        if output_split_sizes is None:44            output = torch.empty_like(input)45        else:46            output = torch.empty(47                size=(sum(output_split_sizes), input.size(1)),48                dtype=input.dtype,49                device=input.device,50            )51        dist.all_to_all_single(52            output,53            input,54            output_split_sizes=output_split_sizes,55            input_split_sizes=input_split_sizes,56            group=group,57        )58        return output59 60    @staticmethod61    def backward(ctx, grad_output):62        return (63            None,64            _AllToAll.apply(65                ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes66            ),67            None,68            None,69        )70 71 72def _all_to_all(73    group: dist.ProcessGroup,74    input: torch.Tensor,75    output_split_sizes: Optional[List[int]],76    input_split_sizes: Optional[List[int]],77) -> torch.Tensor:78    return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)79 80 81def _preprocess(82    expert_mask: torch.Tensor,83    num_experts: int,84    ep_group: dist.ProcessGroup,85) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:86    ep_size = ep_group.size()87    num_local_experts = num_experts // ep_size88    rank = dist.get_rank(ep_group)89    num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))90    input_splits = (91        num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()92    )93    num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)94    output_size = ep_size * num_local_tokens_per_expert_flat.numel()95    num_global_tokens_per_expert_flat = torch.empty(96        output_size,97        dtype=num_local_tokens_per_expert.dtype,98        device=num_local_tokens_per_expert.device,99    )100    dist.all_gather_into_tensor(101        num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group102    )103    num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(104        ep_size, num_local_tokens_per_expert.size(0)105    )106    start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts107    num_global_tokens_per_local_expert = num_global_tokens_per_expert[108        :, start_idx:end_idx109    ].contiguous()110    output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()111    num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(112        dim=0113    ).to(torch.device("cpu"), non_blocking=True)114    num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(115        -1, num_local_experts116    ).to(torch.device("cpu"), non_blocking=True)117    return (118        input_splits,119        output_splits,120        num_global_tokens_per_local_expert,121        num_global_sum_tokens_per_local_expert,122    )123 124 125def _permute(126    tokens: torch.Tensor, routing_map: torch.Tensor127) -> Tuple[torch.Tensor, torch.Tensor]:128    num_tokens, _ = tokens.shape129    num_experts = routing_map.shape[0]130    routing_map = routing_map.bool()131    token_indices = (132        torch.arange(num_tokens, device=routing_map.device)133        .unsqueeze(0)134        .expand(num_experts, -1)135    )136    sorted_indices = token_indices.masked_select(routing_map)137    permuted_input = tokens.index_select(0, sorted_indices)138    return permuted_input, sorted_indices139 140 141def _sort_chunks_by_idxs(142    input: torch.Tensor,143    split_sizes: Union[torch.Tensor, List[int]],144    sorted_idxs: List[int],145) -> torch.Tensor:146    if isinstance(split_sizes, torch.Tensor):147        split_sizes = split_sizes.tolist()148    chunks = torch.split(input, split_sizes, dim=0)149    return torch.cat([chunks[i] for i in sorted_idxs], dim=0)150 151 152def _generate_weights_idx(153    routing_weights: torch.Tensor,154    selected_experts: torch.Tensor,155    num_experts: int,156) -> torch.Tensor:157    num_tokens, topk = routing_weights.shape158    weights_idx = torch.zeros(159        (num_tokens, num_experts),160        dtype=routing_weights.dtype,161        device=routing_weights.device,162    )163    weights_idx.scatter_add_(1, selected_experts, routing_weights)164    return weights_idx165 166 167def _unpermute(168    tokens: torch.Tensor,169    routing_weights: torch.Tensor,170    hidden_states_shape: torch.Size,171    permutation_mapping: torch.Tensor,172    routing_map: torch.Tensor,173) -> torch.Tensor:174    tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())175    tokens = tokens * tokens_weight.unsqueeze(-1)176    hidden_dim = hidden_states_shape[-1]177    unpermuted_tokens = torch.zeros(178        hidden_states_shape, device=tokens.device, dtype=tokens.dtype179    )180    expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)181    unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)182    return unpermuted_tokens183 184 185def token_pre_all2all(186    hidden_states: torch.Tensor,187    expert_mask: torch.Tensor,188    num_experts: int,189    input_splits: List[int],190    output_splits: List[int],191    num_global_tokens_per_local_expert: torch.Tensor,192    group: Optional[dist.ProcessGroup] = None,193) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:194    group = group or dist.group.WORLD195    hidden_dim = hidden_states.size(-1)196    hidden_states = hidden_states.reshape(-1, hidden_dim)197    org_hidden_states_shape = hidden_states.shape198    routing_map = expert_mask.sum(dim=1)199 200    local_permuted_hidden_states, local_input_permutation_mapping = _permute(201        hidden_states, routing_map202    )203    expected_tokens = sum(input_splits)204    actual_tokens = local_permuted_hidden_states.shape[0]205    if expected_tokens != actual_tokens:206        raise RuntimeError(207            f"EP split mismatch: input_splits sum ({expected_tokens}) != "208            f"permuted tokens ({actual_tokens})"209        )210 211    global_permuted_hidden_states = _all_to_all(212        group, local_permuted_hidden_states, output_splits, input_splits213    )214    num_local_experts = num_experts // dist.get_world_size(group)215    permute_order = (216        torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()217    )218    split_sizes = num_global_tokens_per_local_expert.ravel().tolist()219    global_permuted_hidden_states = _sort_chunks_by_idxs(220        global_permuted_hidden_states, split_sizes, permute_order221    )222    return (223        global_permuted_hidden_states,224        routing_map,225        local_input_permutation_mapping,226        org_hidden_states_shape,227    )228 229 230def tokens_post_all2all(231    expert_outputs: torch.Tensor,232    routing_weights: torch.Tensor,233    selected_experts: torch.Tensor,234    num_experts: int,235    input_splits: List[int],236    output_splits: List[int],237    num_global_tokens_per_local_expert: torch.Tensor,238    routing_map: torch.Tensor,239    local_input_permutation_mapping: torch.Tensor,240    org_hidden_states_shape: torch.Size,241    group: Optional[dist.ProcessGroup] = None,242) -> torch.Tensor:243    group = group or dist.group.WORLD244    num_local_experts = num_experts // dist.get_world_size(group)245    unpermute_order = (246        torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()247    )248    split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()249    expert_outputs = _sort_chunks_by_idxs(250        expert_outputs, split_sizes, unpermute_order251    )252    unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)253    weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)254    unpermute_outputs = _unpermute(255        unpermute_outputs,256        weights_idx,257        org_hidden_states_shape,258        local_input_permutation_mapping,259        routing_map,260    )261    return unpermute_outputs262 263 264def expert_forward(265    x: torch.Tensor,266    gate_proj: torch.nn.Linear,267    up_proj: torch.nn.Linear,268    down_proj: torch.nn.Linear,269) -> torch.Tensor:270    gate = torch.nn.functional.silu(gate_proj(x))271    up = up_proj(x)272    return down_proj(gate * up)273 274 275def solution(276    hidden_states: torch.Tensor,277    gate_weight: torch.Tensor,278    gate_bias: Optional[torch.Tensor],279    gate_proj: torch.nn.Linear,280    up_proj: torch.nn.Linear,281    down_proj: torch.nn.Linear,282    num_experts: int,283    top_k: int,284    group: Optional[dist.ProcessGroup] = None,285) -> torch.Tensor:286    if group is None:287        group = _resolve_ep_group_for_narrow_moe(num_experts)288    hidden_dim = hidden_states.size(-1)289    num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)290 291    router_logits = torch.nn.functional.linear(292        hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias293    )294    routing_weights, selected_experts = torch.topk(295        torch.softmax(router_logits, dim=-1), top_k, dim=-1296    )297    expert_mask = torch.nn.functional.one_hot(298        selected_experts, num_classes=num_experts299    ).permute(2, 1, 0)300 301    input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(302        expert_mask, num_experts, group303    )304 305    (306        global_permuted_hidden_states,307        routing_map,308        local_input_permutation_mapping,309        org_hidden_states_shape,310    ) = token_pre_all2all(311        hidden_states,312        expert_mask,313        num_experts,314        input_splits,315        output_splits,316        num_global_tokens_per_local_expert,317        group,318    )319 320    expert_outputs = expert_forward(321        global_permuted_hidden_states, gate_proj, up_proj, down_proj322    )323 324    out = tokens_post_all2all(325        expert_outputs,326        routing_weights,327        selected_experts,328        num_experts,329        input_splits,330        output_splits,331        num_global_tokens_per_local_expert,332        routing_map,333        local_input_permutation_mapping,334        org_hidden_states_shape,335        group,336    )337    return out338