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
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-k
bias1=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 osimport gc+ import functools+ import sysos.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"⋯ 6 unchanged linesgc.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.shapemax_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.shapeD = n2 if k2 == k1 else n2 * 2- if w1.dtype is torch.uint32:- D = D * 8+ if w1.dtype is torch.uint32: D = D * 8key = (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 lineselse: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