Skip to content
KernelIndex
Search⌘K

submission 741375

LunNova · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_triton_ck_v2b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-741375?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
190.8µs
#756 of 782
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:478d4c743d724c03afb2f44fe0e8f46c3d861f10540d612033d6f18e8b39d9d0
license declaredunknown
license concludedunknown
authorsLunNova
imported2026-08-26

Techniques

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

num-warps = 4num_warps=4,
stages = 2num_stages=2,
tile-k = 512BLOCK_K=512,
tile-n = 64BLOCK_N = 64

Kernel source

sub_triton_ck_v2b.py291 lines
"""
sub_triton_ck_v2b.py — Hybrid: Triton stage1 for shapes where it wins, CK stage1 fallback for rest.
Both paths use CK ASM stage2. Shape dispatch via sk = m | (e << 16) | (inter_dim << 32).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl

import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import _moe_sorting_impl, get_block_size_M, use_nt, get_ksplit
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort


# ═══════════════════════════════════════════════════════════════════
# Shapes where Triton stage1 wins (live benchmarks 2026-04-05)
# Key: m | (E << 16) | (inter_dim << 32)
# Default: CK stage1 (faster on 5/7 contest shapes)
# ═══════════════════════════════════════════════════════════════════
_TRITON_STAGE1_SHAPES = {
    128 | (33 << 16) | (512 << 32),    # de512_E32_bs128:  triton 145 vs ck 148 (E=33 incl shared)
    512 | (33 << 16) | (512 << 32),    # de512_E32_bs512:  triton 242 vs ck 252 (E=33 incl shared)
}


# ═══════════════════════════════════════════════════════════════════
# Fused MOE Stage1: Triton GEMM (gate+up) + SiLU + scatter
# Hardcoded config: BLOCK_N=64, BLOCK_K=512, warps=4, stages=2
# (autotuned winner for dep=512, Kp=3584 on gfx950)
# ═══════════════════════════════════════════════════════════════════
@triton.jit
def _moe_stage1_fused(
    a_ptr, w_ptr, a2_ptr,
    a_sc_ptr, w_sc_ptr,
    sorted_ids_ptr, sorted_expert_ids_ptr, num_valid_ids_ptr,
    M, d_expert_pad, K_packed, num_valid_padded,
    stride_a_m, stride_a_k,
    stride_w_e, stride_w_bn,
    stride_as_blk,
    stride_ws_blk, w_sc_expert_offset,
    stride_a2_m, stride_a2_n,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    K_PACKED: tl.constexpr,
    top_k: tl.constexpr,
):
    SCALE_GRP: tl.constexpr = 32
    K_PACK: tl.constexpr = BLOCK_K // 2
    K_SC: tl.constexpr = BLOCK_K // SCALE_GRP
    NUM_K_ITERS: tl.constexpr = (2 * K_PACKED) // BLOCK_K

    pid = tl.program_id(0)
    num_n_blks = tl.cdiv(d_expert_pad, BLOCK_N)
    pid_m = pid // num_n_blks
    pid_n = pid % num_n_blks

    num_valid = tl.load(num_valid_ids_ptr)
    if pid_m * BLOCK_M >= num_valid:
        return

    expert_id = tl.load(sorted_expert_ids_ptr + pid_m)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    valid_m = offs_m < num_valid
    sid = tl.load(sorted_ids_ptr + offs_m, mask=valid_m, other=0)
    token_ids = (sid & 0xFFFFFF)
    topk_ids = sid >> 24
    valid_tok = token_ids < M
    safe_tids = tl.where(valid_tok, token_ids, 0)
    a_mask = valid_m & valid_tok

    offs_k_a = tl.arange(0, K_PACK)
    a_ptrs_base = a_ptr + safe_tids[:, None] * stride_a_m + offs_k_a[None, :] * stride_a_k
    offs_asm = pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)
    offs_ks_flat = tl.arange(0, K_SC * 32)
    a_sc_ptrs_base = a_sc_ptr + offs_asm[:, None] * stride_as_blk + offs_ks_flat[None, :]

    w_base = w_ptr + expert_id * stride_w_e
    offs_k_shuf = tl.arange(0, K_PACK * 16)
    gate_bn_tile = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
    b_gate_ptrs_base = w_base + gate_bn_tile[:, None] * stride_w_bn + offs_k_shuf[None, :]

    ws_base_blk = expert_id * w_sc_expert_offset
    gate_bsn = ws_base_blk + pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
    b_gate_sc_ptrs_base = w_sc_ptr + gate_bsn[:, None] * stride_ws_blk + offs_ks_flat[None, :]

    up_n_offset = d_expert_pad // 16
    up_bn_tile = (up_n_offset + pid_n * (BLOCK_N // 16)) + tl.arange(0, BLOCK_N // 16)
    b_up_ptrs_base = w_base + up_bn_tile[:, None] * stride_w_bn + offs_k_shuf[None, :]

    up_sc_offset = d_expert_pad // 32
    up_bsn = ws_base_blk + up_sc_offset + pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
    b_up_sc_ptrs_base = w_sc_ptr + up_bsn[:, None] * stride_ws_blk + offs_ks_flat[None, :]

    acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    acc_up = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    a_ptrs = a_ptrs_base
    a_sc_ptrs = a_sc_ptrs_base
    b_gate_ptrs = b_gate_ptrs_base
    b_gate_sc_ptrs = b_gate_sc_ptrs_base
    b_up_ptrs = b_up_ptrs_base
    b_up_sc_ptrs = b_up_sc_ptrs_base

    for _ in range(NUM_K_ITERS):
        a = tl.load(a_ptrs, mask=a_mask[:, None], other=0)
        a_sc_raw = tl.load(a_sc_ptrs)
        a_scales = (
            a_sc_raw
            .reshape(BLOCK_M // 32, K_SC // 8, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_M, K_SC)
        )

        b_gate_raw = tl.load(b_gate_ptrs)
        b_gate = (
            b_gate_raw
            .reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_N, K_PACK)
            .trans(1, 0)
        )
        b_gate_sc_raw = tl.load(b_gate_sc_ptrs)
        b_gate_scales = (
            b_gate_sc_raw
            .reshape(BLOCK_N // 32, K_SC // 8, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_N, K_SC)
        )
        acc_gate = tl.dot_scaled(a, a_scales, "e2m1", b_gate, b_gate_scales, "e2m1", acc_gate)

        b_up_raw = tl.load(b_up_ptrs)
        b_up = (
            b_up_raw
            .reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_N, K_PACK)
            .trans(1, 0)
        )
        b_up_sc_raw = tl.load(b_up_sc_ptrs)
        b_up_scales = (
            b_up_sc_raw
            .reshape(BLOCK_N // 32, K_SC // 8, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_N, K_SC)
        )
        acc_up = tl.dot_scaled(a, a_scales, "e2m1", b_up, b_up_scales, "e2m1", acc_up)

        a_ptrs += K_PACK * stride_a_k
        a_sc_ptrs += K_SC * 32
        b_gate_ptrs += K_PACK * 16
        b_gate_sc_ptrs += K_SC * 32
        b_up_ptrs += K_PACK * 16
        b_up_sc_ptrs += K_SC * 32

    result = (acc_gate * tl.sigmoid(acc_gate)) * acc_up

    flat_row = token_ids * top_k + topk_ids
    safe_row = tl.where(a_mask, flat_row, 0)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    out_ptrs = a2_ptr + safe_row[:, None] * stride_a2_m + offs_n[None, :] * stride_a2_n
    out_mask = a_mask[:, None] & (offs_n[None, :] < d_expert_pad)
    tl.store(out_ptrs, result.to(tl.bfloat16), mask=out_mask)


# ═══════════════════════════════════════════════════════════════════
# Entry point
# ═══════════════════════════════════════════════════════════════════
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

    M = hidden_states.shape[0]
    topk = topk_ids.shape[1]
    E = gate_up_weight_shuffled.shape[0]
    N_gateup = gate_up_weight_shuffled.shape[1]
    inter_dim = N_gateup // 2
    model_dim = down_weight_shuffled.shape[1]
    d_hidden_pad = config["d_hidden_pad"]
    d_expert_pad = config["d_expert_pad"]
    K_packed = d_hidden_pad // 2
    K_scale = d_hidden_pad // 32

    dtype = torch.bfloat16
    device = hidden_states.device

    block_m = get_block_size_M(M, topk, E, inter_dim)
    ksplit = get_ksplit(M, topk, E, inter_dim, model_dim)
    non_temporal = use_nt(M, topk, E)

    # ── 1. Token sorting (opus) ──
    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _moe_sorting_impl(
        topk_ids, topk_weights, E, model_dim, dtype, block_m,
        expert_mask=None, num_local_tokens=None, dispatch_policy=0, use_opus=True,
    )
    num_valid_padded = sorted_ids.shape[0]

    # ── 2. Quantize activations (fused with sort → shuffled scales) ──
    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,
    )

    # ── 3. Stage1: dispatch triton vs CK based on shape ──
    sk = M | (E << 16) | (inter_dim << 32)
    use_ck = sk not in _TRITON_STAGE1_SHAPES

    a2 = torch.empty((M, topk, inter_dim), dtype=dtype, device=device)

    if use_ck:
        # CK stage1: gate_up GEMM + SiLU
        w1_scale_e8m0 = gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
        aiter.ck_moe_stage1_fwd(
            a1, gate_up_weight_shuffled, down_weight_shuffled,
            sorted_ids, sorted_expert_ids, num_valid_ids,
            a2, topk, "",
            w1_scale_e8m0, a1_scale, block_m,
            None, QuantType.per_1x32, ActivationType.Silu,
            ksplit, non_temporal, dtype,
        )
    else:
        # Triton fused stage1: GEMM + SiLU + scatter
        a2_flat = a2.view(-1, inter_dim)
        w1_u8 = gate_up_weight_shuffled.view(torch.uint8)
        a1_u8 = a1.view(torch.uint8)
        a1_sc_u8 = a1_scale.view(torch.uint8)
        w1_sc_u8 = gate_up_weight_scale_shuffled.view(torch.uint8)
        stride_w_bn = 16 * K_packed
        stride_as_blk = K_scale * 32
        stride_ws_blk = K_scale * 32
        w_sc_expert_offset = N_gateup // 32

        BLOCK_N = 64
        grid = (triton.cdiv(num_valid_padded, block_m) *
                triton.cdiv(d_expert_pad, BLOCK_N),)

        _moe_stage1_fused[grid](
            a1_u8, w1_u8, a2_flat,
            a1_sc_u8, w1_sc_u8,
            sorted_ids, sorted_expert_ids, num_valid_ids,
            M, d_expert_pad, K_packed, num_valid_padded,
            a1_u8.stride(0), a1_u8.stride(1),
            w1_u8.stride(0), stride_w_bn,
            stride_as_blk,
            stride_ws_blk, w_sc_expert_offset,
            a2_flat.stride(0), a2_flat.stride(1),
            BLOCK_M=block_m,
            BLOCK_N=BLOCK_N,
            BLOCK_K=512,
            K_PACKED=K_packed,
            top_k=topk,
            num_warps=4,
            num_stages=2,
        )

    # ── 4. Quantize intermediate (for CK stage2) ──
    a2_flat = a2.view(-1, inter_dim)
    a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        a2_flat,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=M,
        topk=topk,
        block_size=block_m,
    )
    a2_q = a2_q.view(M, topk, -1)

    # ── 5. CK ASM Stage2: down GEMM ──
    w2_scale_e8m0 = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
    aiter.ck_moe_stage2_fwd(
        a2_q, gate_up_weight_shuffled, down_weight_shuffled,
        sorted_ids, sorted_expert_ids, num_valid_ids,
        moe_buf, topk, "",
        w2_scale_e8m0, a2_scale, block_m,
        sorted_weights, QuantType.per_1x32, ActivationType.Silu,
        non_temporal,
    )

    return moe_buf
scrolls · 291 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