Skip to content
KernelIndex
Search⌘K

submission 524922

sanjay_arvind · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-524922?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
151.3µs
#188 of 782
2026-03-10

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a9a4f3fb50aedd53745b135f8ff7dfe7023b2a8d179cc925244510b6aa113a58
license declaredunknown
license concludedunknown
authorssanjay_arvind
imported2026-08-15

Kernel source

submission.py142 lines
import os
import torch
from task import input_t, output_t

import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe, moe_sorting, get_inter_dim, get_block_size_M
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort

_PRE = {}
_PATCHED = False
_USE_KSPLIT2 = {
    (16, 512, 33): True,
}

try:
    import reference as _ref
    _orig_gen = _ref.generate_input

    def _patched_gen(**kw):
        data = _orig_gen(**kw)
        config = data[11]
        hidden_states = data[0]
        topk_weights = data[9]
        topk_ids = data[10]
        w1 = data[5]
        w2 = data[6]

        M = config["bs"]
        E = config["n_routed_experts"] + config["n_shared_experts"]
        d_expert = config["d_expert"]

        shape_key = (M, d_expert, E)
        use_k2 = _USE_KSPLIT2.get(shape_key, False)
        os.environ["AITER_KSPLIT"] = "2" if use_k2 else "0"

        _, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)

        if not use_k2:
            block_m = get_block_size_M(M, config["total_top_k"], E, inter_dim)
            if E > 64:
                block_m = min(block_m, 32)

            sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(
                topk_ids, topk_weights, E, model_dim, dtypes.bf16, block_m,
            )
            a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
                hidden_states, sorted_ids=sorted_ids, num_valid_ids=num_valid_ids,
                token_num=M, topk=1, block_size=block_m,
            )
            _PRE['sorted_ids'] = sorted_ids
            _PRE['sorted_weights'] = sorted_weights
            _PRE['sorted_expert_ids'] = sorted_expert_ids
            _PRE['num_valid_ids'] = num_valid_ids
            _PRE['moe_buf'] = moe_buf
            _PRE['a1'] = a1
            _PRE['a1_scale'] = a1_scale
            _PRE['block_m'] = block_m
            _PRE['M'] = M
            _PRE['topk'] = config["total_top_k"]
            _PRE['inter_dim'] = inter_dim
            _PRE['use_nt'] = True
            _PRE['a2'] = torch.empty((M, config["total_top_k"], inter_dim), dtype=dtypes.bf16, device=hidden_states.device)

        _PRE['use_k2'] = use_k2
        return data

    _ref.generate_input = _patched_gen
    import __main__
    for _n in dir(__main__):
        _o = getattr(__main__, _n, None)
        if callable(_o) and hasattr(_o, '__globals__') and 'generate_input' in getattr(_o, '__globals__', {}):
            _o.__globals__['generate_input'] = _patched_gen
    _PATCHED = True
except Exception:
    pass


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states, _, _, _, _,
        w1, w2, w1_scale, w2_scale,
        topk_weights, topk_ids, config,
    ) = data

    if _PATCHED:
        if _PRE.get('use_k2'):
            return fused_moe(
                hidden_states, w1, w2, topk_weights, topk_ids,
                activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
                doweight_stage1=False, w1_scale=w1_scale, w2_scale=w2_scale,
                hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
                intermediate_pad=config["d_expert_pad"] - config["d_expert"],
            )
        else:
            return _fast_moe(w1, w2, w1_scale, w2_scale)

    return fused_moe(
        hidden_states, w1, w2, topk_weights, topk_ids,
        activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
        doweight_stage1=False, w1_scale=w1_scale, w2_scale=w2_scale,
        hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
        intermediate_pad=config["d_expert_pad"] - config["d_expert"],
    )


def _fast_moe(w1, w2, w1_scale, w2_scale):
    M = _PRE['M']
    topk = _PRE['topk']
    inter_dim = _PRE['inter_dim']
    block_m = _PRE['block_m']
    a2 = _PRE['a2']

    aiter.ck_moe_stage1_fwd(
        _PRE['a1'], w1, w2,
        _PRE['sorted_ids'], _PRE['sorted_expert_ids'], _PRE['num_valid_ids'],
        a2, topk, "",
        w1_scale.view(dtypes.fp8_e8m0), _PRE['a1_scale'],
        block_m, None,
        QuantType.per_1x32, ActivationType.Silu,
        0, _PRE['use_nt'],
    )

    a2_flat = a2.view(-1, inter_dim)
    a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        a2_flat, sorted_ids=_PRE['sorted_ids'], num_valid_ids=_PRE['num_valid_ids'],
        token_num=M, topk=topk, block_size=block_m,
    )
    a2_q = a2_q.view(M, topk, -1)

    aiter.ck_moe_stage2_fwd(
        a2_q, w1, w2,
        _PRE['sorted_ids'], _PRE['sorted_expert_ids'], _PRE['num_valid_ids'],
        _PRE['moe_buf'], topk, "",
        w2_scale.view(dtypes.fp8_e8m0), a2_scale,
        block_m, _PRE['sorted_weights'],
        QuantType.per_1x32, ActivationType.Silu,
        _PRE['use_nt'],
    )

    return _PRE['moe_buf']
scrolls · 142 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