Skip to content
KernelIndex
Search⌘K

submission 523642

ooousay · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a220900fe36f75bfa5ced70e235fdbc480ac22c04072ce8cd6ef26a96bb2eeab
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.py135 lines
"""
V188: Fix force-NT arg positions + always force NT for all CK Codegen (ksplit=0) shapes.
V176 had wrong arg mapping (args[3]=inter_dim not expert, args[4]=expert not topk).
This caused E=33 M=512 d=2048 (slowest shape) to NOT get force-NT.
"""
import os
import gc
import functools
import sys

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()

# Patch get_2stage_cfgs to force non_temporal_load=True for large-M CK Codegen shapes
_orig_get_2stage_cfgs = get_2stage_cfgs.__wrapped__
@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(*args, **kwargs):
    meta = _orig_get_2stage_cfgs(*args, **kwargs)
    if meta.ksplit == 0:
        # Force non_temporal_load=True for ALL CK Codegen shapes
        # Default heuristic only enables NT for small-M; benchmarks show NT helps large-M too
        if hasattr(meta.stage1, 'keywords') and 'use_non_temporal_load' in meta.stage1.keywords:
            meta.stage1.keywords['use_non_temporal_load'] = True
        if hasattr(meta.stage2, 'keywords') and 'use_non_temporal_load' in meta.stage2.keywords:
            meta.stage2.keywords['use_non_temporal_load'] = True
    return meta
sys.modules['aiter.fused_moe'].get_2stage_cfgs = _patched_get_2stage_cfgs

_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

_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, guw, dw, gus, ds, guw_s, dw_s, gus_s, ds_s, tw, ti, config) = data
    M = hidden_states.shape[0]; E = guw_s.shape[0]
    hp = config["d_hidden_pad"] - config["d_hidden"]
    ip = config["d_expert_pad"] - config["d_expert"]
    de = config["d_expert"]

    if E > 64 and M <= 128: mode = "cktile_e257_k7"
    elif E > 64: mode = "e257_csv"
    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 == "e257_csv":
            os.environ.pop("AITER_KSPLIT", None)
            os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
        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)
        _patched_get_2stage_cfgs.cache_clear(); _current_mode = mode

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

    act = ActivationType.Silu
    if mode == "default" and de <= 512: act = ActivationType.Swiglu

    return fused_moe(hidden_states, guw_s, dw_s, tw, ti, expert_mask=None,
        activation=act, quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=gus_s, w2_scale=ds_s, a1_scale=None, a2_scale=None,
        block_size_M=bsm, hidden_pad=hp, intermediate_pad=ip)
scrolls · 135 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 518255.

"""
- Honest kernel: real GPU work only, no hacks.
- Best config-adaptive fused_moe with optimized ksplit/bsm/activation.
+ V188: Fix force-NT arg positions + always force NT for all CK Codegen (ksplit=0) shapes.
+ V176 had wrong arg mapping (args[3]=inter_dim not expert, args[4]=expert not topk).
+ This caused E=33 M=512 d=2048 (slowest shape) to NOT get force-NT.
"""
import os
import gc
+ import functools
+ import sys
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
⋯ 6 unchanged lines
gc.disable()
- # Pre-allocate sorting buffers to avoid repeated allocation
- _sorting_cache = {}
+ # Patch get_2stage_cfgs to force non_temporal_load=True for large-M CK Codegen shapes
+ _orig_get_2stage_cfgs = get_2stage_cfgs.__wrapped__
+ @functools.lru_cache(maxsize=2048)
+ def _patched_get_2stage_cfgs(*args, **kwargs):
+ meta = _orig_get_2stage_cfgs(*args, **kwargs)
+ if meta.ksplit == 0:
+ # Force non_temporal_load=True for ALL CK Codegen shapes
+ # Default heuristic only enables NT for small-M; benchmarks show NT helps large-M too
+ if hasattr(meta.stage1, 'keywords') and 'use_non_temporal_load' in meta.stage1.keywords:
+ meta.stage1.keywords['use_non_temporal_load'] = True
+ if hasattr(meta.stage2, 'keywords') and 'use_non_temporal_load' in meta.stage2.keywords:
+ meta.stage2.keywords['use_non_temporal_load'] = True
+ return meta
+ sys.modules['aiter.fused_moe'].get_2stage_cfgs = _patched_get_2stage_cfgs
- 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,
- ):
+ _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)
⋯ 9 unchanged lines
)
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,
- )
+ 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,
+ 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,
- ):
+ 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
+ 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,
+ 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,
+ 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,
- )
+ 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)
+ 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
+ (hidden_states, guw, dw, gus, ds, guw_s, dw_s, gus_s, ds_s, tw, ti, config) = data
+ M = hidden_states.shape[0]; E = guw_s.shape[0]
+ hp = config["d_hidden_pad"] - config["d_hidden"]
+ ip = config["d_expert_pad"] - config["d_expert"]
+ de = config["d_expert"]
- 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"]
+ if E > 64 and M <= 128: mode = "cktile_e257_k7"
+ elif E > 64: mode = "e257_csv"
+ elif E <= 64 and M <= 16: mode = "cktile_k7"
+ elif E <= 64 and M <= 128: mode = "cktile_k2"
+ else: mode = "default"
- # 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 == "e257_csv":
+ os.environ.pop("AITER_KSPLIT", None)
+ os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
elif mode == "cktile_k7":
os.environ["AITER_KSPLIT"] = "7"
os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
⋯ 3 unchanged lines
else:
os.environ.pop("AITER_KSPLIT", None)
os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
- get_2stage_cfgs.cache_clear()
- _current_mode = mode
+ _patched_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
+ if mode == "default" and E <= 64: bsm = 64
+ elif mode == "cktile_k2": 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
+ if mode == "default" and de <= 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,
- )
+ return fused_moe(hidden_states, guw_s, dw_s, tw, ti, expert_mask=None,
+ activation=act, quant_type=QuantType.per_1x32, doweight_stage1=False,
+ w1_scale=gus_s, w2_scale=ds_s, a1_scale=None, a2_scale=None,
+ block_size_M=bsm, hidden_pad=hp, intermediate_pad=ip)
scrolls · 217 diff lines total

Best evidence level for this revision: reported

JSON