Skip to content
KernelIndex
Search⌘K

submission 576432

Koren · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 687 lines, June 9 Researcher Reciprocity License v1.0.

submission_2_aiter.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-576432?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
185.1µs
#586 of 782
2026-03-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3abdb224bc85bcade4ce7e7d58e12e4b5a98f719a3d6ca023f33d8b3bc112ba0
license declaredunknown
license concludedunknown
authorsKoren
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4_LOG_PREFIX = "[mxfp4-moe-prof]"

Kernel source

submission_2_aiter.py687 lines
import atexit
from collections import deque
import math
import sys
import time
from typing import Dict, Optional

import torch
from task import input_t, output_t

import aiter
from aiter import ActivationType, QuantType, dtypes, fp4_utils, get_hip_quant
from aiter import get_hip_quant as get_quant
from aiter.fused_moe import get_2stage_cfgs, get_inter_dim, get_padded_M, moe_sorting
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort

from aiter.jit.utils.torch_guard import torch_compile_guard


# Profiler controls (edit these globals directly in this file).
_SUBMISSION_DEBUG = False
_SUBMISSION_SAMPLE_EVERY = 1
_SUBMISSION_EMIT_EVERY = 99
_SUBMISSION_TAIL_SIZE = 64
_LOG_PREFIX = "[mxfp4-moe-prof]"
_STAGE_NAMES = ("sort", "quant1", "stage1", "quant2", "stage2", "two_stage", "e2e")
_CASE_FIELDS = (
    "bs",
    "d_hidden",
    "d_expert",
    "n_routed_experts",
    "n_shared_experts",
    "n_experts_per_token",
    "total_top_k",
    "M",
    "E_runtime",
    "topk_runtime",
    "dtype_hidden",
    "dtype_w1",
)

_DEBUG = bool(_SUBMISSION_DEBUG)
_SAMPLE_EVERY = max(1, int(_SUBMISSION_SAMPLE_EVERY))
_EMIT_EVERY = max(1, int(_SUBMISSION_EMIT_EVERY))
_TAIL_SIZE = max(1, int(_SUBMISSION_TAIL_SIZE))

_stats: Dict[str, object] = {
    "calls_total": 0,
    "calls_timed": 0,
    "by_case": {},
}
_active_profile: Optional[Dict[str, object]] = None


def _log(msg: str) -> None:
    if _DEBUG:
        print(f"{_LOG_PREFIX} {msg}", file=sys.stderr, flush=True)


def _sample_std(total: float, sumsq: float, n: int) -> float:
    if n <= 1:
        return 0.0
    variance = (sumsq - (total * total) / float(n)) / float(n - 1)
    return math.sqrt(max(variance, 0.0))


def _stage_bucket() -> Dict[str, object]:
    return {
        "total": 0.0,
        "sum_sq": 0.0,
        "min": float("inf"),
        "max": 0.0,
        "tail": deque(maxlen=_TAIL_SIZE),
    }


def _new_case_stats() -> Dict[str, object]:
    return {
        "calls_total": 0,
        "calls_timed": 0,
        "gpu_us": {stage: _stage_bucket() for stage in _STAGE_NAMES},
        "cpu_us": {stage: _stage_bucket() for stage in _STAGE_NAMES},
    }


def _cfg_int(config: dict, primary: str, alternate: str, default: int = -1) -> int:
    if primary in config:
        return int(config[primary])
    if alternate in config:
        return int(config[alternate])
    return default


def _make_case_key(config: dict, hidden_states: torch.Tensor, w1: torch.Tensor, topk_ids: torch.Tensor) -> tuple:
    return (
        _cfg_int(config, "bs", "bs"),
        _cfg_int(config, "d_hidden", "dhidden"),
        _cfg_int(config, "d_expert", "dexpert"),
        _cfg_int(config, "n_routed_experts", "nroutedexperts"),
        _cfg_int(config, "n_shared_experts", "nsharedexperts"),
        _cfg_int(config, "n_experts_per_token", "nexpertspertoken"),
        _cfg_int(config, "total_top_k", "totaltopk", int(topk_ids.shape[1])),
        int(hidden_states.shape[0]),
        int(w1.shape[0]),
        int(topk_ids.shape[1]),
        str(hidden_states.dtype),
        str(w1.dtype),
    )


def _format_case_key(case_key: tuple) -> str:
    parts = []
    for idx, field in enumerate(_CASE_FIELDS):
        parts.append(f"{field}={case_key[idx]}")
    return " ".join(parts)


def _ensure_case_stats(case_key: tuple) -> Dict[str, object]:
    by_case = _stats["by_case"]
    if case_key not in by_case:
        by_case[case_key] = _new_case_stats()
    return by_case[case_key]


def _record_case_call(case_key: tuple) -> None:
    _stats["calls_total"] = int(_stats["calls_total"]) + 1
    case_stats = _ensure_case_stats(case_key)
    case_stats["calls_total"] = int(case_stats["calls_total"]) + 1


def _update_stage_bucket(case_stats: Dict[str, object], mode: str, stage: str, value_us: float) -> None:
    bucket = case_stats[mode][stage]
    bucket["total"] += value_us
    bucket["sum_sq"] += value_us * value_us
    bucket["min"] = min(bucket["min"], value_us)
    bucket["max"] = max(bucket["max"], value_us)
    bucket["tail"].append(value_us)


def _begin_profile(case_key: tuple) -> None:
    global _active_profile
    _active_profile = {
        "case_key": case_key,
        "cpu_us": {stage: 0.0 for stage in _STAGE_NAMES},
        "events": [],
    }


def _profile_active() -> bool:
    return _DEBUG and _active_profile is not None


def _time_stage(stage: str, fn):
    if not _profile_active():
        return fn()

    start_event = torch.cuda.Event(enable_timing=True)
    end_event = torch.cuda.Event(enable_timing=True)
    cpu_t0 = time.perf_counter_ns()
    start_event.record()
    out = fn()
    end_event.record()
    cpu_t1 = time.perf_counter_ns()

    _active_profile["cpu_us"][stage] += (cpu_t1 - cpu_t0) / 1000.0
    _active_profile["events"].append((stage, start_event, end_event))
    return out


def _tail_mean_and_std(bucket: Dict[str, object]) -> tuple[float, float]:
    tail = bucket["tail"]
    tail_n = len(tail)
    if tail_n == 0:
        return 0.0, 0.0
    total = float(sum(tail))
    sumsq = float(sum(x * x for x in tail))
    return total / float(tail_n), _sample_std(total, sumsq, tail_n)


def _tail_stats(bucket: Dict[str, object]) -> Dict[str, float]:
    tail = bucket["tail"]
    tail_n = len(tail)
    if tail_n == 0:
        return {"n": 0.0, "avg": 0.0, "std": 0.0}
    avg, std = _tail_mean_and_std(bucket)
    return {
        "n": float(tail_n),
        "avg": avg,
        "std": std,
    }


def _emit_case_line(case_key: tuple, case_stats: Dict[str, object]) -> None:
    if int(case_stats["calls_timed"]) <= 0:
        return

    gpu = case_stats["gpu_us"]
    cpu = case_stats["cpu_us"]
    gpu_e2e = _tail_stats(gpu["e2e"])
    gpu_sort = _tail_stats(gpu["sort"])
    gpu_quant1 = _tail_stats(gpu["quant1"])
    gpu_quant2 = _tail_stats(gpu["quant2"])
    gpu_stage1 = _tail_stats(gpu["stage1"])
    gpu_stage2 = _tail_stats(gpu["stage2"])
    gpu_two_stage = _tail_stats(gpu["two_stage"])

    cpu_e2e = _tail_stats(cpu["e2e"])
    cpu_sort = _tail_stats(cpu["sort"])
    cpu_quant1 = _tail_stats(cpu["quant1"])
    cpu_quant2 = _tail_stats(cpu["quant2"])
    cpu_stage1 = _tail_stats(cpu["stage1"])
    cpu_stage2 = _tail_stats(cpu["stage2"])
    cpu_two_stage = _tail_stats(cpu["two_stage"])

    gpu_quant_avg = gpu_quant1["avg"] + gpu_quant2["avg"]
    cpu_quant_avg = cpu_quant1["avg"] + cpu_quant2["avg"]

    stage_share = {"sort": 0.0, "quant_total": 0.0, "stage1": 0.0, "stage2": 0.0}
    if gpu_e2e["avg"] > 0.0:
        stage_share["sort"] = 100.0 * gpu_sort["avg"] / gpu_e2e["avg"]
        stage_share["quant_total"] = 100.0 * gpu_quant_avg / gpu_e2e["avg"]
        stage_share["stage1"] = 100.0 * gpu_stage1["avg"] / gpu_e2e["avg"]
        stage_share["stage2"] = 100.0 * gpu_stage2["avg"] / gpu_e2e["avg"]

    ranked = sorted(stage_share.items(), key=lambda item: item[1], reverse=True)
    bottleneck, bottleneck_share = ranked[0]
    runner_up, runner_up_share = ranked[1]

    _log(
        f"case[{_format_case_key(case_key)}]\n"
        f"  calls_total={case_stats['calls_total']} calls_timed={case_stats['calls_timed']} "
        f"tail_window_n={int(gpu_e2e['n'])} (all values are tail-window aggregate avg/std in us)\n"
        f"  GPU avg: e2e={gpu_e2e['avg']:.2f}us\n\ttwo_stage={gpu_two_stage['avg']:.2f}us"
        f"\n\tsort={gpu_sort['avg']:.2f}us\n\tquant1={gpu_quant1['avg']:.2f}us quant2={gpu_quant2['avg']:.2f}us "
        f"quant_total={gpu_quant_avg:.2f}us\n\tstage1={gpu_stage1['avg']:.2f}us stage2={gpu_stage2['avg']:.2f}us\n"
        f"\n  GPU std: e2e={gpu_e2e['std']:.2f}us\n\tsort={gpu_sort['std']:.2f}us "
        f"\n\tquant_total={math.sqrt(max(gpu_quant1['std'] * gpu_quant1['std'] + gpu_quant2['std'] * gpu_quant2['std'], 0.0)):.2f}us "
        f"\n\tstage1={gpu_stage1['std']:.2f}us stage2={gpu_stage2['std']:.2f}us\n"
        f"\n\n  CPU avg: e2e={cpu_e2e['avg']:.2f}us\n\ttwo_stage={cpu_two_stage['avg']:.2f}us "
        f"\n\tsort={cpu_sort['avg']:.2f}us\n\tquant1={cpu_quant1['avg']:.2f}us quant2={cpu_quant2['avg']:.2f}us "
        f"quant_total={cpu_quant_avg:.2f}us\n\tstage1={cpu_stage1['avg']:.2f}us stage2={cpu_stage2['avg']:.2f}us\n"
        f"\n  CPU std: e2e={cpu_e2e['std']:.2f}us\n\tsort={cpu_sort['std']:.2f}us "
        f"\n\tquant_total={math.sqrt(max(cpu_quant1['std'] * cpu_quant1['std'] + cpu_quant2['std'] * cpu_quant2['std'], 0.0)):.2f}us "
        f"\n\tstage1={cpu_stage1['std']:.2f}us stage2={cpu_stage2['std']:.2f}us\n"
        f"\n\n  bottleneck_gpu={bottleneck}({bottleneck_share:.1f}% of e2e) "
        f"second={runner_up}({runner_up_share:.1f}% of e2e)"
    )


def _finalize_profile() -> None:
    global _active_profile
    ctx = _active_profile
    if ctx is None:
        return

    gpu_us_by_stage = {stage: 0.0 for stage in _STAGE_NAMES}
    try:
        torch.cuda.synchronize()
        for stage, start_event, end_event in ctx["events"]:
            gpu_us_by_stage[stage] += start_event.elapsed_time(end_event) * 1000.0
    except Exception as err:  # keep profiling best-effort only
        _log(f"profiling synchronize/error: {err}")

    case_key = ctx["case_key"]
    case_stats = _ensure_case_stats(case_key)
    case_stats["calls_timed"] = int(case_stats["calls_timed"]) + 1
    _stats["calls_timed"] = int(_stats["calls_timed"]) + 1

    for stage in _STAGE_NAMES:
        cpu_value = float(ctx["cpu_us"].get(stage, 0.0))
        if cpu_value > 0.0:
            _update_stage_bucket(case_stats, "cpu_us", stage, cpu_value)
        gpu_value = float(gpu_us_by_stage.get(stage, 0.0))
        if gpu_value > 0.0:
            _update_stage_bucket(case_stats, "gpu_us", stage, gpu_value)

    calls_timed = int(case_stats["calls_timed"])
    if calls_timed % _EMIT_EVERY == 0:
        _emit_case_line(case_key, case_stats)
    _active_profile = None


def _final_report() -> None:
    if not _DEBUG:
        return
    if int(_stats["calls_timed"]) == 0:
        _log("no timed calls recorded")
        return
    by_case = _stats["by_case"]
    for case_key, case_stats in sorted(
        by_case.items(),
        key=lambda entry: int(entry[1]["calls_timed"]),
        reverse=True,
    ):
        if int(case_stats["calls_timed"]) > 0:
            _emit_case_line(case_key, case_stats)


atexit.register(_final_report)


def fused_moe_2stages(
    hidden_states,
    w1,  # [expert(local_expert:EP), inter_dim*2, dim] N,K
    w2,  # [expert(local_expert:EP), dim, inter_dim]
    topk,
    sorted_ids,
    sorted_weights,
    sorted_expert_ids,
    num_valid_ids,
    moe_out,
    isG1U1,
    block_size_M,
    activation=ActivationType.Silu,
    quant_type=QuantType.No,
    doweight_stage1=False,
    # following for quant
    q_dtype_a=None,
    q_dtype_w=None,
    w1_scale=None,  # [expert(local_expert:EP), inter_dim, 1]
    w2_scale=None,  # [expert(local_expert:EP), model_dim, 1]
    a1_scale=None,  # [expert(local_expert:EP), 1, model_dim]
    a2_scale=None,  # [expert(local_expert:EP), 1, inter_dim]
    num_local_tokens: Optional[torch.tensor] = None,
    # following for cktile support
    hidden_pad=0,
    intermediate_pad=0,
    bias1=None,
    bias2=None,
    metadata=None,
):
    token_num, _ = hidden_states.shape
    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
    dtype = moe_out.dtype
    device = hidden_states.device
    is_shuffled = getattr(w1, "is_shuffled", False)
    if metadata is None:
        metadata = get_2stage_cfgs(
            get_padded_M(token_num),  # consider token_num > 1024 as prefill
            model_dim,
            inter_dim,
            E,
            topk,
            dtype,
            q_dtype_a,
            q_dtype_w,
            quant_type,
            isG1U1,
            activation,
            doweight_stage1,
            hidden_pad,
            intermediate_pad,
            is_shuffled,
        )

    assert quant_type == QuantType.per_1x32, "only per_1x32 quant type is supported"

    a1, a1_scale = _time_stage(
        "quant1",
        lambda: fused_dynamic_mxfp4_quant_moe_sort(
            hidden_states,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            topk=1,
            block_size=block_size_M,
        ),
    )

    a2 = torch.empty(
        (token_num, topk, inter_dim),
        dtype=dtype,
        device=device,
    )

    extra_stage1_args = {}
    extra_stage2_args = {}

    # print(f"\nRunning stage1 {metadata.stage1.func.__name__}\n", file=sys.stderr, flush=True)
    
    a2 = _time_stage(
        "stage1",
        lambda: metadata.stage1(
            a1,
            w1,
            w2,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            a2,
            topk,
            block_m=block_size_M,
            a1_scale=a1_scale,
            w1_scale=(
                w1_scale.view(dtypes.fp8_e8m0) if w1.dtype == dtypes.fp4x2 else w1_scale
            ),
            sorted_weights=sorted_weights if doweight_stage1 else None,
            **extra_stage1_args,
        ),
    )
    
    a2 = a2.view(-1, inter_dim)
    a2, a2_scale = _time_stage(
        "quant2",
        lambda: fused_dynamic_mxfp4_quant_moe_sort(
            a2,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            topk=topk,
            block_size=block_size_M,
        ),
    )

    a2 = a2.view(token_num, topk, -1)

    # print(f"\nRunning stage2 {metadata.stage2.func.__name__}\n", file=sys.stderr, flush=True)

    _time_stage(
        "stage2",
        lambda: metadata.stage2(
            a2,
            w1,
            w2,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            moe_out,
            topk,
            w2_scale=(
                w2_scale.view(dtypes.fp8_e8m0) if w2.dtype == dtypes.fp4x2 else w2_scale
            ),
            a2_scale=a2_scale,
            block_m=block_size_M,
            sorted_weights=sorted_weights if not doweight_stage1 else None,
            **extra_stage2_args,
        ),
    )

    return moe_out


def fused_moe_fake_submission(
    hidden_states: torch.Tensor,
    w1: torch.Tensor,  # [expert(local_expert:EP), inter_dim*2, dim] N,K
    w2: torch.Tensor,  # [expert(local_expert:EP), dim, inter_dim]
    topk_weight: torch.Tensor,
    topk_ids: torch.Tensor,
    expert_mask: Optional[torch.Tensor] = None,  # EP
    activation: int = ActivationType.Silu.value,
    quant_type: int = QuantType.No.value,
    doweight_stage1: bool = False,
    # following for quant
    w1_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), inter_dim, 1]
    w2_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), model_dim, 1]
    a1_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), 1, model_dim]
    a2_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), 1, inter_dim]
    # following for tuning
    block_size_M: int = -1,
    num_local_tokens: Optional[torch.Tensor] = None,
    moe_sorting_dispatch_policy: bool = 0,
    dtype: Optional[torch.dtype] = None,
    hidden_pad: int = 0,
    intermediate_pad: int = 0,
    bias1: Optional[torch.Tensor] = None,
    bias2: Optional[torch.Tensor] = None,
) -> torch.Tensor:
    device = topk_ids.device
    M, topk = topk_ids.shape
    dtype = hidden_states.dtype if dtype is None else dtype
    model_dim = w2.shape[1]
    moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
    return moe_buf


quant_remap = {QuantType.per_128x128: QuantType.per_1x128}


@torch_compile_guard(gen_fake=fused_moe_fake_submission)
def fused_moe_submission(
    hidden_states: torch.Tensor,
    w1: torch.Tensor,  # [expert(local_expert:EP), inter_dim*2, dim] N,K
    w2: torch.Tensor,  # [expert(local_expert:EP), dim, inter_dim]
    topk_weight: torch.Tensor,
    topk_ids: torch.Tensor,
    expert_mask: Optional[torch.Tensor] = None,  # EP
    activation: int = ActivationType.Silu.value,
    quant_type: int = QuantType.No.value,
    doweight_stage1: bool = False,
    # following for quant
    w1_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), inter_dim, 1]
    w2_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), model_dim, 1]
    a1_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), 1, model_dim]
    a2_scale: Optional[torch.Tensor] = None,  # [expert(local_expert:EP), 1, inter_dim]
    # following for tuning
    block_size_M: int = -1,
    num_local_tokens: Optional[torch.Tensor] = None,
    moe_sorting_dispatch_policy: bool = 0,
    dtype: Optional[torch.dtype] = None,
    hidden_pad: int = 0,
    intermediate_pad: int = 0,
    bias1: Optional[torch.Tensor] = None,
    bias2: Optional[torch.Tensor] = None,
) -> torch.Tensor:
    # We do such convert since custom_op schema restriction on block_size_M, and Enum type
    activation = ActivationType(activation)
    quant_type = QuantType(quant_type)
    if block_size_M == -1:
        block_size_M = None
    """user API"""
    M, topk = topk_ids.shape
    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)

    assert w1.shape[1] in [
        inter_dim,
        inter_dim * 2,
    ], f"Invalid MoE weight: {w1.shape=} {w2.shape=}"
    isG1U1 = inter_dim != w1.shape[1]
    isShuffled = getattr(w1, "is_shuffled", False)

    global_E = E
    if expert_mask is not None:
        global_E = expert_mask.numel()
    dtype = hidden_states.dtype if dtype is None else dtype
    assert dtype in [
        dtypes.fp16,
        dtypes.bf16,
    ], f"Fused_moe unsupported out dtype: {dtype}"
    quant_type = quant_remap.get(quant_type, quant_type)
    q_dtype_w = w1.dtype
    q_dtype_a = w1.dtype if w1.dtype != torch.uint32 else dtypes.fp8
    bf16_fp8_bound = 512
    if quant_type == QuantType.per_1x32:
        if activation == ActivationType.Swiglu:
            if get_gfx() != "gfx950" or M < bf16_fp8_bound:
                q_dtype_a = dtypes.bf16
            elif M >= bf16_fp8_bound:
                q_dtype_a = dtypes.fp8
        else:
            q_dtype_a = dtypes.fp4x2

    metadata = get_2stage_cfgs(
        get_padded_M(M),  # consider token_num > 1024 as prefill
        model_dim,
        inter_dim,
        E,
        topk,
        dtype,
        q_dtype_a,
        q_dtype_w,
        quant_type,
        isG1U1,
        activation,
        doweight_stage1,
        hidden_pad,
        intermediate_pad,
        isShuffled,
    )

    block_size_M = metadata.block_m if block_size_M is None else block_size_M
    # Ensure block_size_M is int (metadata.block_m from CSV may be float)
    if block_size_M is not None:
        block_size_M = int(block_size_M)

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _time_stage(
        "sort",
        lambda: moe_sorting(
            topk_ids,
            topk_weight,
            global_E,
            model_dim,
            dtype,
            block_size_M,
            expert_mask,
            num_local_tokens,
            moe_sorting_dispatch_policy,
        ),
    )

    assert not metadata.run_1stage, "Expected 2 stage kernel"
    return _time_stage(
        "two_stage",
        lambda: fused_moe_2stages(
            hidden_states,
            w1,
            w2,
            topk,
            sorted_ids,
            sorted_weights,
            sorted_expert_ids,
            num_valid_ids,
            moe_buf,
            isG1U1,
            block_size_M,
            activation=activation,
            quant_type=quant_type,
            doweight_stage1=doweight_stage1,
            q_dtype_a=q_dtype_a,
            q_dtype_w=q_dtype_w,
            w1_scale=w1_scale,
            w2_scale=w2_scale,
            a1_scale=a1_scale,
            a2_scale=a2_scale,
            num_local_tokens=num_local_tokens,
            # following for cktile support
            hidden_pad=hidden_pad,
            intermediate_pad=intermediate_pad,
            bias1=bias1,
            bias2=bias2,
            metadata=metadata,
        )
    )


def custom_kernel(data: input_t) -> output_t:
    """
    Submission template for DeepSeek-R1 MXFP4 MoE kernel.

    Input data tuple:
        hidden_states:                [M, d_hidden]                           bf16
        gate_up_weight:               [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (raw)
        down_weight:                  [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (raw)
        gate_up_weight_scale:         [E, 2*d_expert_pad, scale_K]            e8m0   (raw)
        down_weight_scale:            [E, d_hidden_pad, scale_K]              e8m0   (raw)
        gate_up_weight_shuffled:      [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (shuffled)
        down_weight_shuffled:         [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (shuffled)
        gate_up_weight_scale_shuffled:[padded, flat]                          e8m0   (shuffled)
        down_weight_scale_shuffled:   [padded, flat]                          e8m0   (shuffled)
        topk_weights:                 [M, total_top_k]                        float32
        topk_ids:                     [M, total_top_k]                        int32
        config:                       dict

    Returns:
        output: [M, d_hidden] bf16
    """
    (
        hidden_states,
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    case_key = _make_case_key(config, hidden_states, gate_up_weight_shuffled, topk_ids)
    _record_case_call(case_key)
    should_time = _DEBUG and (int(_stats["calls_total"]) % _SAMPLE_EVERY == 0)

    if should_time:
        _begin_profile(case_key)

    try:
        output = _time_stage(
            "e2e",
            lambda: fused_moe_submission(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                topk_weights,
                topk_ids,
                expert_mask=None,
                activation=ActivationType.Silu,
                quant_type=QuantType.per_1x32,
                doweight_stage1=False,
                w1_scale=gate_up_weight_scale_shuffled,
                w2_scale=down_weight_scale_shuffled,
                a1_scale=None,
                a2_scale=None,
                hidden_pad=hidden_pad,
                intermediate_pad=intermediate_pad,
                block_size_M=-1,
            ),
        )
    finally:
        if should_time:
            _finalize_profile()

    return output
scrolls · 687 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON