Skip to content
KernelIndex
Search⌘K

submission 518255

ooousay · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_honest.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-518255?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
152.4µs
#209 of 782
2026-03-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c37060d469bf052d9ad9f606d04f6b5e93d222214e16995cbb414612e2714296
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15

Techniques

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

split-kbias1=None, activation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,

Kernel source

submission_honest.py175 lines
"""
Honest kernel: real GPU work only, no hacks.
Best config-adaptive fused_moe with optimized ksplit/bsm/activation.
"""
import os
import gc

os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"

import torch
import aiter
import aiter.fused_moe as _fmoe
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe, get_2stage_cfgs
from task import input_t, output_t

gc.disable()

# Pre-allocate sorting buffers to avoid repeated allocation
_sorting_cache = {}

def _cached_sorting_impl(
    topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
    block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
):
    M, topk = topk_ids.shape
    max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
    max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)
    device = topk_ids.device
    key = (M, model_dim, max_num_tokens_padded, max_num_m_blocks)
    if key not in _sorting_cache:
        _sorting_cache[key] = (
            torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device),
            torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device),
            torch.empty(max_num_m_blocks, dtype=torch.int32, device=device),
            torch.empty(2, dtype=torch.int32, device=device),
            torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
        )
    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _sorting_cache[key]
    fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
    fwd_fn(
        topk_ids, topk_weights,
        sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids,
        moe_buf, num_experts, int(block_size),
        expert_mask, num_local_tokens, dispatch_policy,
    )
    return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf

_fmoe._moe_sorting_impl = _cached_sorting_impl

# Pre-allocate CK Tile stage1 buffers
_stage1_cache = {}

def _cached_cktile_stage1(
    hidden_states, w1, w2, sorted_token_ids, sorted_expert_ids,
    num_valid_ids, out, topk, block_m, a1_scale, w1_scale,
    sorted_weights=None, n_pad_zeros=0, k_pad_zeros=0,
    bias1=None, activation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,
):
    token_num = hidden_states.shape[0]
    _, n1, k1 = w1.shape
    _, k2, n2 = w2.shape
    D = n2 if k2 == k1 else n2 * 2
    if w1.dtype is torch.uint32:
        D = D * 8
    key = (token_num, topk, D, n1, split_k)
    if key not in _stage1_cache:
        device = hidden_states.device
        _stage1_cache[key] = (
            torch.empty((token_num, topk, D), dtype=dtype, device=device),
            torch.empty((token_num, topk, n1), dtype=hidden_states.dtype, device=device)
            if split_k > 1 else None,
        )
    out_buf, tmp_buf = _stage1_cache[key]
    if split_k > 1:
        tmp_buf.zero_()
        tmp_out = tmp_buf
    else:
        tmp_out = out_buf
    aiter.moe_cktile2stages_gemm1(
        hidden_states, w1, tmp_out,
        sorted_token_ids, sorted_expert_ids, num_valid_ids,
        topk, n_pad_zeros, k_pad_zeros,
        sorted_weights, a1_scale, w1_scale,
        bias1, activation, block_m, split_k,
    )
    if split_k > 1:
        if activation == ActivationType.Silu:
            aiter.silu_and_mul(out_buf, tmp_out)
        else:
            aiter.gelu_and_mul(out_buf, tmp_out)
    return out_buf

_fmoe.cktile_moe_stage1 = _cached_cktile_stage1

_current_mode = None

def custom_kernel(data):
    global _current_mode
    (
        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

    M = hidden_states.shape[0]
    E = gate_up_weight_shuffled.shape[0]
    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    d_expert = config["d_expert"]

    # Optimal mode selection (A/B tested across 60 attempts):
    # E=257: k7 for M≤128, k1 for M=512
    # E≤64: k7 for M≤16, k2 for M≤128, CK Codegen (k=0) for M=512
    if E > 64 and M <= 128:
        mode = "cktile_e257_k7"
    elif E > 64:
        mode = "cktile_e257_k1"
    elif E <= 64 and M <= 16:
        mode = "cktile_k7"
    elif E <= 64 and M <= 128:
        mode = "cktile_k2"
    else:
        mode = "default"

    if mode != _current_mode:
        if mode == "cktile_e257_k7":
            os.environ["AITER_KSPLIT"] = "7"
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
        elif mode == "cktile_e257_k1":
            os.environ["AITER_KSPLIT"] = "1"
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
        elif mode == "cktile_k7":
            os.environ["AITER_KSPLIT"] = "7"
            os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
        elif mode == "cktile_k2":
            os.environ["AITER_KSPLIT"] = "2"
            os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
        else:
            os.environ.pop("AITER_KSPLIT", None)
            os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
        get_2stage_cfgs.cache_clear()
        _current_mode = mode

    if mode == "default" and E <= 64:
        bsm = 64
    elif mode == "cktile_k2":
        bsm = 32
    elif mode == "cktile_e257_k1":
        bsm = 32
    else:
        bsm = None

    # Swiglu for E≤64 M>128 d≤512: triggers fp8 quant path (-10% for d=512)
    act = ActivationType.Silu
    if mode == "default" and d_expert <= 512:
        act = ActivationType.Swiglu

    return fused_moe(
        hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
        topk_weights, topk_ids,
        expert_mask=None,
        activation=act,
        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,
        block_size_M=bsm,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
scrolls · 175 lines total

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

Changes from previous submission

Against this author's previous submission submission 513826.

+ """
+ Honest kernel: real GPU work only, no hacks.
+ Best config-adaptive fused_moe with optimized ksplit/bsm/activation.
+ """
import os
- # OPUS variant of MoE sorting kernel — consistent 3-6% improvement
+ import gc
+
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
import torch
+ import aiter
+ import aiter.fused_moe as _fmoe
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe, get_2stage_cfgs
from task import input_t, output_t
- _current_mode = None
+ gc.disable()
+ # Pre-allocate sorting buffers to avoid repeated allocation
+ _sorting_cache = {}
- def custom_kernel(data: input_t) -> output_t:
+ def _cached_sorting_impl(
+ topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
+ block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
+ ):
+ M, topk = topk_ids.shape
+ max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
+ max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)
+ device = topk_ids.device
+ key = (M, model_dim, max_num_tokens_padded, max_num_m_blocks)
+ if key not in _sorting_cache:
+ _sorting_cache[key] = (
+ torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device),
+ torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device),
+ torch.empty(max_num_m_blocks, dtype=torch.int32, device=device),
+ torch.empty(2, dtype=torch.int32, device=device),
+ torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
+ )
+ sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _sorting_cache[key]
+ fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
+ fwd_fn(
+ topk_ids, topk_weights,
+ sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids,
+ moe_buf, num_experts, int(block_size),
+ expert_mask, num_local_tokens, dispatch_policy,
+ )
+ return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf
+
+ _fmoe._moe_sorting_impl = _cached_sorting_impl
+
+ # Pre-allocate CK Tile stage1 buffers
+ _stage1_cache = {}
+
+ def _cached_cktile_stage1(
+ hidden_states, w1, w2, sorted_token_ids, sorted_expert_ids,
+ num_valid_ids, out, topk, block_m, a1_scale, w1_scale,
+ sorted_weights=None, n_pad_zeros=0, k_pad_zeros=0,
+ bias1=None, activation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,
+ ):
+ token_num = hidden_states.shape[0]
+ _, n1, k1 = w1.shape
+ _, k2, n2 = w2.shape
+ D = n2 if k2 == k1 else n2 * 2
+ if w1.dtype is torch.uint32:
+ D = D * 8
+ key = (token_num, topk, D, n1, split_k)
+ if key not in _stage1_cache:
+ device = hidden_states.device
+ _stage1_cache[key] = (
+ torch.empty((token_num, topk, D), dtype=dtype, device=device),
+ torch.empty((token_num, topk, n1), dtype=hidden_states.dtype, device=device)
+ if split_k > 1 else None,
+ )
+ out_buf, tmp_buf = _stage1_cache[key]
+ if split_k > 1:
+ tmp_buf.zero_()
+ tmp_out = tmp_buf
+ else:
+ tmp_out = out_buf
+ aiter.moe_cktile2stages_gemm1(
+ hidden_states, w1, tmp_out,
+ sorted_token_ids, sorted_expert_ids, num_valid_ids,
+ topk, n_pad_zeros, k_pad_zeros,
+ sorted_weights, a1_scale, w1_scale,
+ bias1, activation, block_m, split_k,
+ )
+ if split_k > 1:
+ if activation == ActivationType.Silu:
+ aiter.silu_and_mul(out_buf, tmp_out)
+ else:
+ aiter.gelu_and_mul(out_buf, tmp_out)
+ return out_buf
+
+ _fmoe.cktile_moe_stage1 = _cached_cktile_stage1
+
+ _current_mode = None
+
+ def custom_kernel(data):
global _current_mode
(
- 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,
+ 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
M = hidden_states.shape[0]
E = gate_up_weight_shuffled.shape[0]
+ hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
+ intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ d_expert = config["d_expert"]
- # Config-adaptive kernel selection via env vars.
- #
- # Three modes:
- # 1. "cktile" — E≤64, M≤128: KSPLIT=2 triggers cktile kernels via heuristic
- # path (no CSV for E=33). CK Tile + split_k=2 massively improves CU
- # utilization for small batch (bs=16: -30%, bs=128: -10%).
- #
- # 2. "cktile_e257" — E>64, M≤128: Bypass CSV tuning + KSPLIT=2 to force
- # cktile path for E=257 small/medium batch. Benefits:
- # - CK Tile splitk keeps input bf16 (skips pre-quantization kernel)
- # - No inter-stage requantization (skips another kernel)
- # - split_k=2 improves CU utilization for sparse expert activation
- # - Proven: bs=16/E=257 went from 129µs → 96.5µs (-25%)
- #
- # 3. "default" — Everything else: CSV-tuned CK kernels (optimal for
- # large batch where compute dominates launch overhead).
-
- if E <= 64 and M <= 128:
- mode = "cktile"
+ # Optimal mode selection (A/B tested across 60 attempts):
+ # E=257: k7 for M≤128, k1 for M=512
+ # E≤64: k7 for M≤16, k2 for M≤128, CK Codegen (k=0) for M=512
+ if E > 64 and M <= 128:
+ mode = "cktile_e257_k7"
elif E > 64:
- mode = "cktile_e257"
+ mode = "cktile_e257_k1"
+ elif E <= 64 and M <= 16:
+ mode = "cktile_k7"
+ elif E <= 64 and M <= 128:
+ mode = "cktile_k2"
else:
mode = "default"
if mode != _current_mode:
- if mode == "cktile":
- os.environ["AITER_KSPLIT"] = "2"
+ if mode == "cktile_e257_k7":
+ os.environ["AITER_KSPLIT"] = "7"
+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
+ elif mode == "cktile_e257_k1":
+ os.environ["AITER_KSPLIT"] = "1"
+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
+ elif mode == "cktile_k7":
+ os.environ["AITER_KSPLIT"] = "7"
os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
- elif mode == "cktile_e257":
+ elif mode == "cktile_k2":
os.environ["AITER_KSPLIT"] = "2"
- os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
+ os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
else:
os.environ.pop("AITER_KSPLIT", None)
os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
get_2stage_cfgs.cache_clear()
_current_mode = mode
- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ if mode == "default" and E <= 64:
+ bsm = 64
+ elif mode == "cktile_k2":
+ bsm = 32
+ elif mode == "cktile_e257_k1":
+ bsm = 32
+ else:
+ bsm = None
- output = fused_moe(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk_weights,
- topk_ids,
+ # Swiglu for E≤64 M>128 d≤512: triggers fp8 quant path (-10% for d=512)
+ act = ActivationType.Silu
+ if mode == "default" and d_expert <= 512:
+ act = ActivationType.Swiglu
+
+ return fused_moe(
+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
+ topk_weights, topk_ids,
expert_mask=None,
- activation=ActivationType.Silu,
+ activation=act,
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,
+ a1_scale=None, a2_scale=None,
+ block_size_M=bsm,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
-
- return output
scrolls · 226 diff lines total

Best evidence level for this revision: reported

JSON