Skip to content
KernelIndex
Search⌘K

submission 611420

Aniket Sadashiva · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v503.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-611420?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
130.0µs
#87 of 782
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:414d9a95897a0a36af49273b691afb5f941a6e8b0a5529706aa5dda12605afa0
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15

Techniques

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

fp4a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
split-kactivation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,

Kernel source

submission_v503.py460 lines
"""v503"""
import functools
import os
import sys

import torch

from dataclasses import replace

_dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
_flydsl_s3_stage2 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
try:
    with open(_dsv3_path, "r") as f:
        lines = f.readlines()
    header = lines[0].strip()
    modified_lines = [header + "\n"]
    for line in lines[1:]:
        stripped = line.strip()
        if not stripped:
            continue
        fields = stripped.split(",")
        try:
            token_val = int(fields[1])
            expert_val = int(fields[4])
        except (ValueError, IndexError):
            modified_lines.append(line)
            continue
        if expert_val == 257 and token_val <= 128:
            fields[14] = "2"
            modified_lines.append(",".join(fields) + "\n")
        elif expert_val == 257 and token_val == 512:
            flydsl_fields = list(fields)
            flydsl_fields[19] = _flydsl_s3_stage2
            flydsl_fields[20] = "0.1%"
            if len(flydsl_fields) > 25:
                flydsl_fields[25] = ""
            modified_lines.append(",".join(flydsl_fields) + "\n")
            fallback_fields = list(fields)
            if len(fallback_fields) > 25:
                fallback_fields[25] = "flydsl_fallback"
            else:
                fallback_fields.append("flydsl_fallback")
            modified_lines.append(",".join(fallback_fields) + "\n")
        else:
            modified_lines.append(line)
    with open(_dsv3_path, "w") as f:
        f.writelines(modified_lines)
except Exception as e:
    print(f"[v490] dsv3 error: {e}", file=sys.stderr)

_csv_header = (
    "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,"
    "q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,"
    "err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag"
)
_common = (
    "ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,"
    "torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
)
_k1_512 = (
    "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_512 = (
    "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_k1_2048 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_2048 = (
    "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_k2_512_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
_csv_rows = [
    f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,90.0,{_k2_512_flydsl},0.1%,219.79,0,781.78,2884.18,",
    f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,0,{_k2_512},0.0%,129.79,0,781.78,2884.18,flydsl_fallback",
    f"256,512,7168,2048,33,9,{_common},64,0,0,{_k1_2048},0.0%,180.0,{_k2_2048_flydsl},0.1%,455.08,0,1475.47,5323.27,",
    f"256,512,7168,2048,33,9,{_common},128,0,0,{_k1_2048},0.0%,0,{_k2_2048},0.0%,275.08,0,1475.47,5323.27,flydsl_fallback",
]
try:
    e33_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
    with open(e33_path, "w") as f:
        f.write(_csv_header + "\n")
        for row in _csv_rows:
            f.write(row + "\n")
except Exception:
    pass

os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["AITER_USE_NT"] = "1"
os.environ.pop("FLIR_CK_LDS128", None)
os.environ["FLIR_MOE_STAGE1_SCHED"] = "1"
os.environ["FLIR_MOE_STAGE2_SCHED"] = "1"
os.environ["FLIR_MOE_STAGE2_PERSIST_M"] = "1"

from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2


_SORT_BUFS = {}


def _cached_moe_sorting_impl(
    topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
    block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
):
    device = topk_ids.device
    M, topk = topk_ids.shape
    key = (M, num_experts, block_size, model_dim)

    if key not in _SORT_BUFS:
        max_num_tokens_padded = int(M * topk + num_experts * block_size - topk)
        max_num_m_blocks = int(
            (max_num_tokens_padded + block_size - 1) // block_size
        )
        _SORT_BUFS[key] = (
            torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
            torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
            torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
            torch.empty(2, dtype=dtypes.i32, device=device),
            torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
        )

    sid, sw, sei, nvi, mb = _SORT_BUFS[key]
    fwd = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
    fwd(
        topk_ids, topk_weights, sid, sw, sei, nvi, mb,
        num_experts, int(block_size), expert_mask, num_local_tokens,
        dispatch_policy,
    )
    return sid, sw, sei, nvi, mb


fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl


_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs

_CKTILE_BUFS = {}


def _cached_cktile_moe_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

    buf_key = (token_num, topk, D, w1.shape[1], split_k, hidden_states.device)
    if buf_key not in _CKTILE_BUFS:
        _CKTILE_BUFS[buf_key] = (
            torch.empty((token_num, topk, D), dtype=dtype, device=hidden_states.device),
            torch.zeros(
                (token_num, topk, w1.shape[1]), dtype=hidden_states.dtype,
                device=hidden_states.device,
            ) if split_k > 1 else None,
        )

    out_buf, tmp_buf = _CKTILE_BUFS[buf_key]

    if split_k > 1:
        tmp_buf.zero_()
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, tmp_buf,
            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,
        )
        aiter.silu_and_mul(out_buf, tmp_buf)
    else:
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, out_buf,
            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,
        )
    return out_buf


def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
    return fused_moe_mod.MOEMetadata(
        functools.partial(
            _cached_cktile_moe_stage1,
            n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
            k_pad_zeros=hidden_pad // 128 * 128,
            activation=activation,
            split_k=split_k,
        ),
        functools.partial(
            fused_moe_mod.cktile_moe_stage2,
            n_pad_zeros=hidden_pad // 64 * 64,
            k_pad_zeros=intermediate_pad // 128 * 128,
            activation=activation,
        ),
        16, split_k, False, False, True,
    )


@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(
    token, model_dim, inter_dim, expert, topk, dtype,
    q_dtype_a, q_dtype_w, q_type, use_g1u1, activation,
    doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True,
):
    common = (
        model_dim == 7168 and topk == 9
        and dtype == dtypes.bf16 and q_dtype_a == dtypes.fp4x2
        and q_dtype_w == dtypes.fp4x2 and q_type == QuantType.per_1x32
        and use_g1u1 and not doweight_stage1 and is_shuffled
    )

    if not common:
        return _ORIGINAL_GET_2STAGE_CFGS(
            token, model_dim, inter_dim, expert, topk, dtype,
            q_dtype_a, q_dtype_w, q_type, use_g1u1, activation,
            doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
        )

    if expert == 257 and inter_dim == 256 and token == 16:
        return _make_cktile_metadata(
            hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
        )
    if expert == 257 and inter_dim == 256 and token == 128:
        return _make_cktile_metadata(
            hidden_pad, intermediate_pad, use_g1u1, activation, split_k=4,
        )

    if expert == 33 and inter_dim == 512 and token <= 128:
        return _make_cktile_metadata(
            hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
        )

    return _ORIGINAL_GET_2STAGE_CFGS(
        token, model_dim, inter_dim, expert, topk, dtype,
        q_dtype_a, q_dtype_w, q_type, use_g1u1, activation,
        doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
    )


fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs


_DIRECT_BUFS = {}


def _get_split_k(E, M, inter_dim):
    if E == 257 and inter_dim == 256:
        if M == 16:
            return 2
        if M == 128:
            return 4
    if E == 33 and inter_dim == 512:
        if M == 16:
            return 2
        if M == 128:
            return 1
    return 0


def _direct_cktile_pipeline(
    hidden_states, w1, w2, w1_scale, w2_scale,
    topk_ids, topk_weights, E, M, topk, model_dim,
    inter_dim, hidden_pad, intermediate_pad, split_k,
):
    block_m = 16

    sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
        topk_ids, topk_weights, E, model_dim, torch.bfloat16,
        block_m, None, None, 0, True,
    )

    _, 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

    n_pad1 = intermediate_pad // 64 * 64 * 2
    k_pad1 = hidden_pad // 128 * 128
    n_pad2 = hidden_pad // 64 * 64
    k_pad2 = intermediate_pad // 128 * 128

    buf_key = (M, topk, D, w1.shape[1], split_k)
    if buf_key not in _DIRECT_BUFS:
        dev = hidden_states.device
        out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=dev)
        tmp = (
            torch.zeros((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=dev)
            if split_k > 1 else None
        )
        _DIRECT_BUFS[buf_key] = (out, tmp)

    out_buf, tmp_buf = _DIRECT_BUFS[buf_key]

    w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
    w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)

    if split_k > 1:
        tmp_buf.zero_()
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, tmp_buf,
            sid, sei, nvi,
            topk, n_pad1, k_pad1,
            None, None, w1_scale_e8m0, None,
            ActivationType.Silu, block_m, split_k,
        )
        aiter.silu_and_mul(out_buf, tmp_buf)
    else:
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, out_buf,
            sid, sei, nvi,
            topk, n_pad1, k_pad1,
            None, None, w1_scale_e8m0, None,
            ActivationType.Silu, block_m, 1,
        )

    aiter.moe_cktile2stages_gemm2(
        out_buf, w2, mb,
        sid, sei, nvi,
        topk, n_pad2, k_pad2,
        sw, None, w2_scale_e8m0, None,
        ActivationType.Silu, block_m,
    )

    return mb


_A2_BUFS = {}

_S3_CK_K1 = (
    "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)

_DIRECT_CK_CONFIGS = {
    (257, 256): (_S3_CK_K1, 32, 16, 256, 256, "reduce"),
    (33, 512): (_k1_512, 32, 16, 256, 256, "reduce"),
    (33, 2048): (_k1_2048, 64, 16, 256, 256, "atomic"),
}


def _direct_ck_flydsl_pipeline(
    hidden_states, w1, w2, w1_scale, w2_scale,
    topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
):
    ck_k1, block_m, fly_tm, fly_tn, fly_tk, fly_mode = _DIRECT_CK_CONFIGS[(E, inter_dim)]

    sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
        topk_ids, topk_weights, E, model_dim, torch.bfloat16,
        block_m, None, None, 0, True,
    )

    a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
        hidden_states, sorted_ids=sid, num_valid_ids=nvi,
        token_num=M, topk=1, block_size=block_m,
    )

    buf_key = (M, topk, inter_dim)
    if buf_key not in _A2_BUFS:
        _A2_BUFS[buf_key] = torch.empty(
            (M, topk, inter_dim), dtype=torch.bfloat16, device=hidden_states.device,
        )
    a2 = _A2_BUFS[buf_key]

    w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
    aiter.ck_moe_stage1_fwd(
        a1, w1, w2, sid, sei, nvi, a2, topk,
        ck_k1, w1_scale_e8m0, a1_scale, block_m,
        None, QuantType.per_1x32, ActivationType.Silu, 0, True,
        torch.bfloat16,
    )

    a2_flat = a2.view(-1, inter_dim)
    a2_quant, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        a2_flat, sorted_ids=sid, num_valid_ids=nvi,
        token_num=M, topk=topk, block_size=block_m,
    )
    a2_quant = a2_quant.view(M, topk, -1)

    w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
    flydsl_moe_stage2(
        inter_states=a2_quant, w2=w2,
        sorted_token_ids=sid, sorted_expert_ids=sei,
        num_valid_ids=nvi, out=mb, topk=topk,
        tile_m=fly_tm, tile_n=fly_tn, tile_k=fly_tk,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
        mode=fly_mode,
        w2_scale=w2_scale_e8m0, a2_scale=a2_scale,
        sorted_weights=sw,
    )

    return mb


def custom_kernel(data: input_t) -> output_t:
    (
        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"]

    M = hidden_states.shape[0]
    E = gate_up_weight.shape[0]
    inter_dim = config["d_expert"]
    model_dim = hidden_states.shape[1]
    topk = topk_ids.shape[1]

    split_k = _get_split_k(E, M, inter_dim)

    if split_k > 0 and model_dim == 7168 and topk == 9:
        return _direct_cktile_pipeline(
            hidden_states,
            gate_up_weight_shuffled, down_weight_shuffled,
            gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
            topk_ids, topk_weights, E, M, topk, model_dim,
            inter_dim, hidden_pad, intermediate_pad, split_k,
        )

    if M == 512 and model_dim == 7168 and topk == 9 and (E, inter_dim) in _DIRECT_CK_CONFIGS:
        return _direct_ck_flydsl_pipeline(
            hidden_states,
            gate_up_weight_shuffled, down_weight_shuffled,
            gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
            topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
        )

    return fused_moe_mod.fused_moe(
        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,
    )
scrolls · 460 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 601016.

- """v367: v363 + enable dormant FlyDSL stage1 scheduler calls for S3.
-
- - dsv3 CSV: ksplit=2 for E=257 bs<=128, FlyDSL t64x256x256_reduce for S3
- - CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch
- - E=33 CSV: S6 FlyDSL reduce, S7 FlyDSL t64x128x256_atomic
- - OPUS=1 globally
- - NT=1 globally
- - Cached sorting buffers
- - Direct `moe_cktile2stages_gemm1/2` path for S1/S2 only
- """
+ """v503"""
import functools
import os
import sys
⋯ 2 unchanged lines
from dataclasses import replace
- # ── Step 1: Modify dsv3 CSV for E=257: ksplit=2 on bs<=128, FlyDSL stage2 on bs=512 ──
_dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
_flydsl_s3_stage2 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
try:
⋯ 13 unchanged lines
modified_lines.append(line)
continue
if expert_val == 257 and token_val <= 128:
- fields[14] = "2" # ksplit=2
+ fields[14] = "2"
modified_lines.append(",".join(fields) + "\n")
elif expert_val == 257 and token_val == 512:
- # FlyDSL reduce stage2 for S3, with CK fallback
flydsl_fields = list(fields)
flydsl_fields[19] = _flydsl_s3_stage2
flydsl_fields[20] = "0.1%"
if len(flydsl_fields) > 25:
flydsl_fields[25] = ""
modified_lines.append(",".join(flydsl_fields) + "\n")
- # CK fallback row
fallback_fields = list(fields)
if len(fallback_fields) > 25:
fallback_fields[25] = "flydsl_fallback"
⋯ 5 unchanged lines
with open(_dsv3_path, "w") as f:
f.writelines(modified_lines)
except Exception as e:
- print(f"[v329] dsv3 error: {e}", file=sys.stderr)
+ print(f"[v490] dsv3 error: {e}", file=sys.stderr)
- # ── Step 2: E=33 CSV for S4/S5 ksplit=2, S6/S7 tuned configs ──
_csv_header = (
"cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,"
"q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,"
⋯ 3 unchanged lines
"ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,"
"torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
)
- _k1_small = (
- "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
- "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- )
- _k2_small = (
- "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
- "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
- )
_k1_512 = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
⋯ 10 unchanged lines
"moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
- _k2_512_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_reduce"
- _k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic"
+ _k2_512_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
+ _k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
_csv_rows = [
- # S4 (bs=16, d=512): ksplit=2
- f"256,16,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,47,0,67,7691,",
- # S5 (bs=128, d=512): ksplit=2
- f"256,128,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,58,0,434,6263,",
- # S6 (bs=512, d=512): probe FlyDSL reduce stage2 behind a tagged CK fallback
f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,90.0,{_k2_512_flydsl},0.1%,219.79,0,781.78,2884.18,",
f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,0,{_k2_512},0.0%,129.79,0,781.78,2884.18,flydsl_fallback",
- # S7 (bs=512, d=2048): probe FlyDSL atomic stage2 behind a tagged CK fallback
f"256,512,7168,2048,33,9,{_common},64,0,0,{_k1_2048},0.0%,180.0,{_k2_2048_flydsl},0.1%,455.08,0,1475.47,5323.27,",
f"256,512,7168,2048,33,9,{_common},128,0,0,{_k1_2048},0.0%,0,{_k2_2048},0.0%,275.08,0,1475.47,5323.27,flydsl_fallback",
]
⋯ 6 unchanged lines
except Exception:
pass
- # ── Step 3: Global OPUS=1, NT=1 (no per-shape switching!) ──
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["AITER_USE_NT"] = "1"
+ os.environ.pop("FLIR_CK_LDS128", None)
+ os.environ["FLIR_MOE_STAGE1_SCHED"] = "1"
+ os.environ["FLIR_MOE_STAGE2_SCHED"] = "1"
+ os.environ["FLIR_MOE_STAGE2_PERSIST_M"] = "1"
- # ── Step 4: Repair the known FlyDSL stage1 FP4 source bug on the runner
- # and enable its dormant hot-loop scheduler calls. ──
- _flydsl_stage1_bug_path = (
- "/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py"
- )
- try:
- with open(_flydsl_stage1_bug_path, "r") as f:
- _flydsl_src = f.read()
- _old_stage1_sig = (
- "def compute_f8f6f4_tile(\n"
- " acc_gate_in,\n"
- " acc_up_in,\n"
- " b_tile_in_gate,\n"
- " b_tile_in_up,\n"
- " lds_base,\n"
- " *,\n"
- " a0_prefetch=None,\n"
- " a_scale=None,\n"
- " b_scale_gate=None,\n"
- " b_scale_up=None,\n"
- " prefetch_epilogue: bool = False,\n"
- " ):"
- )
- _new_stage1_sig = (
- "def compute_f8f6f4_tile(\n"
- " acc_gate_in,\n"
- " acc_up_in,\n"
- " b_tile_in,\n"
- " lds_base,\n"
- " *,\n"
- " a0_prefetch=None,\n"
- " a_scale=None,\n"
- " b_scale=None,\n"
- " prefetch_epilogue: bool = False,\n"
- " ):"
- )
- if _old_stage1_sig in _flydsl_src:
- _flydsl_src = _flydsl_src.replace(_old_stage1_sig, _new_stage1_sig)
- _stage1_sched_disabled = (
- " # hot_loop_scheduler()\n"
- " gpu.barrier()"
- )
- _stage1_sched_enabled = (
- " hot_loop_scheduler()\n"
- " gpu.barrier()"
- )
- if _stage1_sched_disabled in _flydsl_src:
- _flydsl_src = _flydsl_src.replace(
- _stage1_sched_disabled,
- _stage1_sched_enabled,
- 3,
- )
- with open(_flydsl_stage1_bug_path, "w") as f:
- f.write(_flydsl_src)
- except Exception as e:
- print(f"[v367] flydsl stage1 source patch skipped: {e}", file=sys.stderr)
-
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod
- import aiter.ops.flydsl
- from aiter.ops.flydsl.utils import is_flydsl_available as _is_flydsl_available
+ from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
- print(f"[v329] flydsl available: {_is_flydsl_available()}", file=sys.stderr)
-
- # ═══════════════════════════════════════════════════════════════════════
- # FlyDSL stage1 wrapper (matches CK stage1 signature)
- # ═══════════════════════════════════════════════════════════════════════
- def _flydsl_stage1_wrapper(
- hidden_states,
- w1,
- w2, # unused by stage1
- sorted_token_ids,
- sorted_expert_ids,
- num_valid_ids,
- out,
- topk,
- block_m=32,
- a1_scale=None,
- w1_scale=None,
- kernelName="",
- sorted_weights=None,
- tile_m=32,
- tile_n=256,
- tile_k=256,
- a_dtype="fp4",
- b_dtype="fp4",
- out_dtype="bf16",
- **_kwargs,
- ):
- return aiter.ops.flydsl.flydsl_moe_stage1(
- a=hidden_states,
- w1=w1,
- sorted_token_ids=sorted_token_ids,
- sorted_expert_ids=sorted_expert_ids,
- num_valid_ids=num_valid_ids,
- out=out,
- topk=topk,
- tile_m=tile_m,
- tile_n=tile_n,
- tile_k=tile_k,
- a_dtype=a_dtype,
- b_dtype=b_dtype,
- out_dtype=out_dtype,
- w1_scale=w1_scale,
- a1_scale=a1_scale,
- sorted_weights=sorted_weights,
- )
-
-
- # ═══════════════════════════════════════════════════════════════════════
- # Cached sorting buffers (zero alloc after warmup)
- # ═══════════════════════════════════════════════════════════════════════
_SORT_BUFS = {}
⋯ 31 unchanged lines
fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl
- # ═══════════════════════════════════════════════════════════════════════
- # Direct CKTile path for official S1/S2 only
- # ═══════════════════════════════════════════════════════════════════════
- _DIRECT_SMALL_DISABLED = False
- _CKTILE_DIRECT_BUFS = {}
+ _ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
+ _CKTILE_BUFS = {}
- def _is_direct_small_shape(config, topk):
- total_experts = config["n_routed_experts"] + config["n_shared_experts"]
- return (
- config["d_hidden"] == 7168
- and topk == 9
- and total_experts == 257
- and config["d_expert"] == 256
- and config["bs"] in (16, 128)
- )
-
- def _get_direct_cktile_workspace(hidden_states, w1, w2, topk, total_experts, block_m):
+ def _cached_cktile_moe_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
- model_dim = hidden_states.shape[1]
- inter_dim = n2 if k2 == k1 else n2 * 2
+ D = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
- inter_dim *= 8
+ D = D * 8
- token_num = hidden_states.shape[0]
- max_num_tokens_padded = int(token_num * topk + total_experts * block_m - topk)
- max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
- key = (
- hidden_states.device,
- hidden_states.dtype,
- token_num,
- topk,
- total_experts,
- model_dim,
- n1,
- inter_dim,
- block_m,
- )
- workspace = _CKTILE_DIRECT_BUFS.get(key)
- if workspace is None:
- workspace = {
- "sorted_ids": torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=hidden_states.device),
- "sorted_weights": torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=hidden_states.device),
- "sorted_expert_ids": torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=hidden_states.device),
- "num_valid_ids": torch.empty(2, dtype=dtypes.i32, device=hidden_states.device),
- "moe_buf": torch.empty((token_num, model_dim), dtype=hidden_states.dtype, device=hidden_states.device),
- "tmp_out": torch.empty((token_num, topk, n1), dtype=hidden_states.dtype, device=hidden_states.device),
- "stage1_out": torch.empty((token_num, topk, inter_dim), dtype=hidden_states.dtype, device=hidden_states.device),
- }
- _CKTILE_DIRECT_BUFS[key] = workspace
- return workspace
+ buf_key = (token_num, topk, D, w1.shape[1], split_k, hidden_states.device)
+ if buf_key not in _CKTILE_BUFS:
+ _CKTILE_BUFS[buf_key] = (
+ torch.empty((token_num, topk, D), dtype=dtype, device=hidden_states.device),
+ torch.zeros(
+ (token_num, topk, w1.shape[1]), dtype=hidden_states.dtype,
+ device=hidden_states.device,
+ ) if split_k > 1 else None,
+ )
+ out_buf, tmp_buf = _CKTILE_BUFS[buf_key]
- def _direct_small_cktile(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- gate_up_weight_scale_shuffled,
- down_weight_scale_shuffled,
- topk_weights,
- topk_ids,
- config,
- ):
- block_m = 16
- split_k = 2
- topk = topk_ids.shape[1]
- total_experts = config["n_routed_experts"] + config["n_shared_experts"]
- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
- workspace = _get_direct_cktile_workspace(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk,
- total_experts,
- block_m,
- )
+ if split_k > 1:
+ tmp_buf.zero_()
+ aiter.moe_cktile2stages_gemm1(
+ hidden_states, w1, tmp_buf,
+ 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,
+ )
+ aiter.silu_and_mul(out_buf, tmp_buf)
+ else:
+ aiter.moe_cktile2stages_gemm1(
+ hidden_states, w1, out_buf,
+ 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,
+ )
+ return out_buf
- sorted_ids = workspace["sorted_ids"]
- sorted_weights = workspace["sorted_weights"]
- sorted_expert_ids = workspace["sorted_expert_ids"]
- num_valid_ids = workspace["num_valid_ids"]
- moe_buf = workspace["moe_buf"]
- tmp_out = workspace["tmp_out"]
- stage1_out = workspace["stage1_out"]
- aiter.moe_sorting_fwd(
- topk_ids,
- topk_weights,
- sorted_ids,
- sorted_weights,
- sorted_expert_ids,
- num_valid_ids,
- moe_buf,
- total_experts,
- block_m,
- None,
- None,
- 0,
- )
-
- tmp_out.zero_()
- aiter.moe_cktile2stages_gemm1(
- hidden_states,
- gate_up_weight_shuffled,
- tmp_out,
- sorted_ids,
- sorted_expert_ids,
- num_valid_ids,
- topk,
- intermediate_pad // 64 * 64 * 2,
- hidden_pad // 128 * 128,
- None,
- None,
- gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),
- None,
- ActivationType.Silu,
- block_m,
- split_k,
- )
- aiter.silu_and_mul(stage1_out, tmp_out)
-
- aiter.moe_cktile2stages_gemm2(
- stage1_out,
- down_weight_shuffled,
- moe_buf,
- sorted_ids,
- sorted_expert_ids,
- num_valid_ids,
- topk,
- hidden_pad // 64 * 64,
- intermediate_pad // 128 * 128,
- sorted_weights,
- None,
- down_weight_scale_shuffled.view(dtypes.fp8_e8m0),
- None,
- ActivationType.Silu,
- block_m,
- 1,
- )
- return moe_buf
-
-
- # ═══════════════════════════════════════════════════════════════════════
- # Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, FlyDSL stage1 for bs=512
- # ═══════════════════════════════════════════════════════════════════════
- _ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
-
-
def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
return fused_moe_mod.MOEMetadata(
functools.partial(
- fused_moe_mod.cktile_moe_stage1,
+ _cached_cktile_moe_stage1,
n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
k_pad_zeros=hidden_pad // 128 * 128,
activation=activation,
⋯ 29 unchanged lines
doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
)
- # E=257 bs<512 → CKTile split_k=2 (quant-skip)
- if expert == 257 and inter_dim == 256 and token < 512:
+ if expert == 257 and inter_dim == 256 and token == 16:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
)
+ if expert == 257 and inter_dim == 256 and token == 128:
+ return _make_cktile_metadata(
+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=4,
+ )
- # E=33 d=512 bs<=128 → CKTile split_k=2 (quant-skip)
if expert == 33 and inter_dim == 512 and token <= 128:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
)
- default = _ORIGINAL_GET_2STAGE_CFGS(
+ return _ORIGINAL_GET_2STAGE_CFGS(
token, model_dim, inter_dim, expert, topk, dtype,
q_dtype_a, q_dtype_w, q_type, use_g1u1, activation,
doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
)
- if (
- token == 512
- and expert == 257
- and inter_dim == 256
- and _is_flydsl_available()
- ):
- default = replace(
- default,
- stage1=functools.partial(
- _flydsl_stage1_wrapper,
- tile_m=128,
- tile_n=256,
- tile_k=256,
- a_dtype="fp4",
- b_dtype="fp4",
- out_dtype="bf16",
- ),
+
+ fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
+
+
+ _DIRECT_BUFS = {}
+
+
+ def _get_split_k(E, M, inter_dim):
+ if E == 257 and inter_dim == 256:
+ if M == 16:
+ return 2
+ if M == 128:
+ return 4
+ if E == 33 and inter_dim == 512:
+ if M == 16:
+ return 2
+ if M == 128:
+ return 1
+ return 0
+
+
+ def _direct_cktile_pipeline(
+ hidden_states, w1, w2, w1_scale, w2_scale,
+ topk_ids, topk_weights, E, M, topk, model_dim,
+ inter_dim, hidden_pad, intermediate_pad, split_k,
+ ):
+ block_m = 16
+
+ sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
+ topk_ids, topk_weights, E, model_dim, torch.bfloat16,
+ block_m, None, None, 0, True,
+ )
+
+ _, 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
+
+ n_pad1 = intermediate_pad // 64 * 64 * 2
+ k_pad1 = hidden_pad // 128 * 128
+ n_pad2 = hidden_pad // 64 * 64
+ k_pad2 = intermediate_pad // 128 * 128
+
+ buf_key = (M, topk, D, w1.shape[1], split_k)
+ if buf_key not in _DIRECT_BUFS:
+ dev = hidden_states.device
+ out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=dev)
+ tmp = (
+ torch.zeros((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=dev)
+ if split_k > 1 else None
)
+ _DIRECT_BUFS[buf_key] = (out, tmp)
- # S7 (E=33, d=2048, bs=512): block_m=64 + NT
- if expert == 33 and inter_dim == 2048 and token == 512:
- return replace(default, block_m=64, use_non_temporal_load=True)
+ out_buf, tmp_buf = _DIRECT_BUFS[buf_key]
- return default
+ w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
+ w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
+ if split_k > 1:
+ tmp_buf.zero_()
+ aiter.moe_cktile2stages_gemm1(
+ hidden_states, w1, tmp_buf,
+ sid, sei, nvi,
+ topk, n_pad1, k_pad1,
+ None, None, w1_scale_e8m0, None,
+ ActivationType.Silu, block_m, split_k,
+ )
+ aiter.silu_and_mul(out_buf, tmp_buf)
+ else:
+ aiter.moe_cktile2stages_gemm1(
+ hidden_states, w1, out_buf,
+ sid, sei, nvi,
+ topk, n_pad1, k_pad1,
+ None, None, w1_scale_e8m0, None,
+ ActivationType.Silu, block_m, 1,
+ )
- fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
+ aiter.moe_cktile2stages_gemm2(
+ out_buf, w2, mb,
+ sid, sei, nvi,
+ topk, n_pad2, k_pad2,
+ sw, None, w2_scale_e8m0, None,
+ ActivationType.Silu, block_m,
+ )
+ return mb
- # ═══════════════════════════════════════════════════════════════════════
- # Dispatch — simple, no state switching
- # ═══════════════════════════════════════════════════════════════════════
- def custom_kernel(data: input_t) -> output_t:
- global _DIRECT_SMALL_DISABLED
+ _A2_BUFS = {}
+
+ _S3_CK_K1 = (
+ "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
+ "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ )
+
+ _DIRECT_CK_CONFIGS = {
+ (257, 256): (_S3_CK_K1, 32, 16, 256, 256, "reduce"),
+ (33, 512): (_k1_512, 32, 16, 256, 256, "reduce"),
+ (33, 2048): (_k1_2048, 64, 16, 256, 256, "atomic"),
+ }
+
+
+ def _direct_ck_flydsl_pipeline(
+ hidden_states, w1, w2, w1_scale, w2_scale,
+ topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
+ ):
+ ck_k1, block_m, fly_tm, fly_tn, fly_tk, fly_mode = _DIRECT_CK_CONFIGS[(E, inter_dim)]
+
+ sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
+ topk_ids, topk_weights, E, model_dim, torch.bfloat16,
+ block_m, None, None, 0, True,
+ )
+
+ a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ hidden_states, sorted_ids=sid, num_valid_ids=nvi,
+ token_num=M, topk=1, block_size=block_m,
+ )
+
+ buf_key = (M, topk, inter_dim)
+ if buf_key not in _A2_BUFS:
+ _A2_BUFS[buf_key] = torch.empty(
+ (M, topk, inter_dim), dtype=torch.bfloat16, device=hidden_states.device,
+ )
+ a2 = _A2_BUFS[buf_key]
+
+ w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
+ aiter.ck_moe_stage1_fwd(
+ a1, w1, w2, sid, sei, nvi, a2, topk,
+ ck_k1, w1_scale_e8m0, a1_scale, block_m,
+ None, QuantType.per_1x32, ActivationType.Silu, 0, True,
+ torch.bfloat16,
+ )
+
+ a2_flat = a2.view(-1, inter_dim)
+ a2_quant, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ a2_flat, sorted_ids=sid, num_valid_ids=nvi,
+ token_num=M, topk=topk, block_size=block_m,
+ )
+ a2_quant = a2_quant.view(M, topk, -1)
+
+ w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
+ flydsl_moe_stage2(
+ inter_states=a2_quant, w2=w2,
+ sorted_token_ids=sid, sorted_expert_ids=sei,
+ num_valid_ids=nvi, out=mb, topk=topk,
+ tile_m=fly_tm, tile_n=fly_tn, tile_k=fly_tk,
+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
+ mode=fly_mode,
+ w2_scale=w2_scale_e8m0, a2_scale=a2_scale,
+ sorted_weights=sw,
+ )
+
+ return mb
+
+
+ def custom_kernel(data: input_t) -> output_t:
(
hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
⋯ 5 unchanged lines
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
- if not _DIRECT_SMALL_DISABLED and _is_direct_small_shape(config, topk_ids.shape[1]):
- try:
- return _direct_small_cktile(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- gate_up_weight_scale_shuffled,
- down_weight_scale_shuffled,
- topk_weights,
- topk_ids,
- config,
- )
- except Exception as e:
- _DIRECT_SMALL_DISABLED = True
- print(f"[v350] direct S1/S2 path disabled: {e}", file=sys.stderr)
+ M = hidden_states.shape[0]
+ E = gate_up_weight.shape[0]
+ inter_dim = config["d_expert"]
+ model_dim = hidden_states.shape[1]
+ topk = topk_ids.shape[1]
+ split_k = _get_split_k(E, M, inter_dim)
+
+ if split_k > 0 and model_dim == 7168 and topk == 9:
+ return _direct_cktile_pipeline(
+ hidden_states,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
+ topk_ids, topk_weights, E, M, topk, model_dim,
+ inter_dim, hidden_pad, intermediate_pad, split_k,
+ )
+
+ if M == 512 and model_dim == 7168 and topk == 9 and (E, inter_dim) in _DIRECT_CK_CONFIGS:
+ return _direct_ck_flydsl_pipeline(
+ hidden_states,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
+ topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
+ )
+
return fused_moe_mod.fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids,
scrolls · 686 diff lines total

Best evidence level for this revision: reported

JSON