Skip to content
KernelIndex
Search⌘K

submission 754196

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v329.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754196?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
133.6µs
#93 of 782
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5111877487356be8c8076cbab4ce0ac2801de6d89a302513cf2ddf14f4738fae
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

fp4fp4 = torch.empty((M, dep // 2), dtype=torch.uint8, device=gemm_out.device)
fused-epiloguedef _fused_silu_mul_quant_kernel(
num-warps = 4num_warps=4,
stages = 3num_stages=3,
tile-m = 128BLOCK_M=128, QB=32,
tile-n = 256SR_BN = 256 if M >= 64 else 128

Kernel source

submission_v329.py1141 lines
"""
v329: v268 + eliminate dhp padding waste.
dh=7168 is already divisible by 1024/512/256/128/64/32, so next_power_of_2
padding to 8192 wastes 12.5% of:
- GEMM1 K-iterations (8->7)
- GEMM2 N-tiles (64->56 for BN2=128)
- Quantization grid (256->224 scale groups)
- Buffer allocations
Weights keep their dhp-based shapes (pre-padded by harness), kernels just
access the first dh-relevant portion.
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl


@triton.jit
def _dynamic_mxfp4_quant_kernel(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m, stride_x_n,
    stride_x_fp4_m, stride_x_fp4_n,
    stride_bs_m, stride_bs_n,
    M: tl.constexpr, N: tl.constexpr,
    scaleN: tl.constexpr,
    scaleM_pad: tl.constexpr,
    scaleN_pad: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SCALING_MODE: tl.constexpr,
    SHUFFLE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    stride_x_m = tl.cast(stride_x_m, tl.int64)
    stride_x_n = tl.cast(stride_x_n, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
    x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
    x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
    x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
    x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
    amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.log2(amax).floor() - 2
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
    quant_scale = tl.exp2(-scale_e8m0_unbiased)
    qx = x * quant_scale
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    e = (qx >> 23) & 0xFF
    m = qx & 0x7FFFFF
    E8_BIAS: tl.constexpr = 127
    E2_BIAS: tl.constexpr = 1
    adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
    m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
    e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
    e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
    e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
    e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
    evens, odds = tl.split(e2m1_value)
    out_tensor = evens | (odds << 4)
    out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)
    out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
    out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
    tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".cg")
    bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    bs_offs_n = pid_n
    bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
    bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
    tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask, cache_modifier=".cg")


def dynamic_mxfp4_quant(x):
    M, N = x.shape
    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
    scaleN_valid = triton.cdiv(N, 32)
    scaleN = triton.cdiv(scaleN_valid, 8) * 8
    blockscale = torch.empty(
        (triton.cdiv(M, 256) * 256, scaleN),
        dtype=torch.uint8, device=x.device,
    )
    grid = (triton.cdiv(M, 128), scaleN_valid)
    _dynamic_mxfp4_quant_kernel[grid](
        x, x_fp4, blockscale,
        *x.stride(), *x_fp4.stride(), *blockscale.stride(),
        M=M, N=N, scaleN=scaleN_valid,
        scaleM_pad=triton.cdiv(M, 32) * 32,
        scaleN_pad=scaleN,
        BLOCK_SIZE=128, MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0, SHUFFLE=False,
    )
    blockscale = blockscale[:M, :scaleN_valid].contiguous()
    return (x_fp4, blockscale)


@triton.jit
def _fused_silu_mul_quant_kernel(
    gemm_out_ptr, fp4_ptr, scale_ptr,
    dep,
    stride_gm, stride_gn,
    stride_fp4_m, stride_fp4_n,
    stride_sc_m, stride_sc_n,
    M,
    BLOCK_M: tl.constexpr,
    QB: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    stride_gm = tl.cast(stride_gm, tl.int64)
    stride_gn = tl.cast(stride_gn, tl.int64)
    stride_fp4_m = tl.cast(stride_fp4_m, tl.int64)
    stride_fp4_n = tl.cast(stride_fp4_n, tl.int64)
    m_offs = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    n_offs = pid_n * QB + tl.arange(0, QB)
    m_mask = m_offs < M
    n_mask = n_offs < dep
    gate_offs = m_offs[:, None] * stride_gm + n_offs[None, :] * stride_gn
    up_offs = m_offs[:, None] * stride_gm + (n_offs[None, :] + dep) * stride_gn
    mask = m_mask[:, None] & n_mask[None, :]
    gate = tl.load(gemm_out_ptr + gate_offs, mask=mask, other=0).to(tl.float32)
    up = tl.load(gemm_out_ptr + up_offs, mask=mask, other=0).to(tl.float32)
    x = (gate * tl.sigmoid(gate) * up).to(tl.bfloat16).to(tl.float32)
    amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_unb = tl.log2(amax).floor() - 2
    scale_unb = tl.clamp(scale_unb, min=-127, max=127)
    quant_scale = tl.exp2(-scale_unb)
    qx = x * quant_scale
    bs_e8m0 = scale_unb.to(tl.uint8) + 127
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    e = (qx >> 23) & 0xFF
    m = qx & 0x7FFFFF
    E8_BIAS: tl.constexpr = 127
    E2_BIAS: tl.constexpr = 1
    adj = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
    m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adj, m)
    e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
    e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
    e2m1_val = ((s >> 28) | e2m1_tmp).to(tl.uint8)
    e2m1_val = tl.reshape(e2m1_val, [BLOCK_M, QB // 2, 2])
    evens, odds = tl.split(e2m1_val)
    packed = evens | (odds << 4)
    out_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    out_n = pid_n * QB // 2 + tl.arange(0, QB // 2)
    out_offs = out_m[:, None] * stride_fp4_m + out_n[None, :] * stride_fp4_n
    out_mask = (out_m < M)[:, None] & (out_n < (dep // 2))[None, :]
    tl.store(fp4_ptr + out_offs, packed, mask=out_mask, cache_modifier=".cg")
    sc_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    sc_offs = sc_m[:, None] * stride_sc_m + pid_n * stride_sc_n
    tl.store(scale_ptr + sc_offs, bs_e8m0, mask=(sc_m < M)[:, None], cache_modifier=".cg")


def fused_silu_mul_quant(gemm_out, dep):
    M = gemm_out.shape[0]
    fp4 = torch.empty((M, dep // 2), dtype=torch.uint8, device=gemm_out.device)
    scaleN = triton.cdiv(dep, 32)
    scale = torch.empty((M, scaleN), dtype=torch.uint8, device=gemm_out.device)
    grid = (triton.cdiv(M, 128), scaleN)
    _fused_silu_mul_quant_kernel[grid](
        gemm_out, fp4, scale,
        dep,
        gemm_out.stride(0), gemm_out.stride(1),
        fp4.stride(0), fp4.stride(1),
        scale.stride(0), scale.stride(1),
        M,
        BLOCK_M=128, QB=32,
    )
    return fp4, scale


@triton.jit
def _remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    tall_xcds = GRID_MN % NUM_XCDS
    tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    new_pid = tl.where(
        xcd < tall_xcds,
        xcd * pids_per_xcd + local_pid,
        tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid,
    )
    return new_pid


# ---- Triton counting sort kernels ----

@triton.jit
def _moe_count_and_offsets_kernel(
    topk_ids_ptr, counts_ptr, offsets_ptr, cum_blocks_ptr,
    done_ptr, total, E, num_ctas,
    BM: tl.constexpr, BLOCK: tl.constexpr, BLOCK_E: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < total
    ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
    tl.atomic_add(counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
    old_done = tl.atomic_add(done_ptr, 1)
    if old_done == num_ctas - 1:
        idx = tl.arange(0, BLOCK_E)
        e_mask = idx < E
        counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
        token_offsets = tl.cumsum(counts, axis=0)
        tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_mask)
        n_blocks = tl.where(counts > 0, (counts + BM - 1) // BM, tl.zeros([BLOCK_E], dtype=tl.int64))
        block_offsets = tl.cumsum(n_blocks, axis=0)
        tl.store(cum_blocks_ptr + 1 + idx, block_offsets, mask=e_mask)


@triton.jit
def _moe_scatter_kernel(
    topk_ids_ptr, topk_weights_ptr,
    sorted_token_idx_ptr, sorted_weights_ptr,
    offsets_ptr, write_counts_ptr,
    reverse_idx_ptr, token_counts_ptr,
    topk, total,
    BLOCK: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < total
    ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
    weights = tl.load(topk_weights_ptr + offs, mask=mask, other=0.0)
    token_idx = (offs // topk).to(tl.int32)
    old = tl.atomic_add(write_counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
    base = tl.load(offsets_ptr + ids, mask=mask, other=0).to(tl.int32)
    write_pos = (base + old).to(tl.int64)
    tl.store(sorted_token_idx_ptr + write_pos, token_idx, mask=mask)
    tl.store(sorted_weights_ptr + write_pos, weights, mask=mask)
    tk_old = tl.atomic_add(token_counts_ptr + token_idx, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
    rev_pos = token_idx.to(tl.int64) * topk + tk_old.to(tl.int64)
    tl.store(reverse_idx_ptr + rev_pos, write_pos.to(tl.int32), mask=mask)


@triton.jit
def _single_cta_sort_kernel(
    topk_ids_ptr, topk_weights_ptr, sorted_token_idx_ptr, sorted_weights_ptr,
    counts_ptr, offsets_ptr, cum_blocks_ptr, write_counts_ptr,
    reverse_idx_ptr, token_counts_ptr,
    total, E, M, topk,
    BM: tl.constexpr, BLOCK: tl.constexpr, BLOCK_E: tl.constexpr,
):
    idx = tl.arange(0, BLOCK_E)
    e_mask = idx < E
    tl.store(counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
    tl.store(write_counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
    e1_mask = idx < (E + 1)
    tl.store(offsets_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
    tl.store(cum_blocks_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
    offs = tl.arange(0, BLOCK)
    t_mask = offs < M
    tl.store(token_counts_ptr + offs, tl.zeros([BLOCK], dtype=tl.int32), mask=t_mask)

    tl.debug_barrier()

    mask = offs < total
    ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
    weights = tl.load(topk_weights_ptr + offs, mask=mask, other=0.0)
    tl.atomic_add(counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)

    tl.debug_barrier()

    counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
    token_offsets = tl.cumsum(counts, axis=0)
    tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_mask)
    n_blocks = tl.where(counts > 0, (counts + BM - 1) // BM, tl.zeros([BLOCK_E], dtype=tl.int64))
    block_offsets = tl.cumsum(n_blocks, axis=0)
    tl.store(cum_blocks_ptr + 1 + idx, block_offsets, mask=e_mask)

    tl.debug_barrier()

    token_idx = (offs // topk).to(tl.int32)
    old = tl.atomic_add(write_counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
    base = tl.load(offsets_ptr + ids, mask=mask, other=0).to(tl.int32)
    write_pos = (base + old).to(tl.int64)
    tl.store(sorted_token_idx_ptr + write_pos, token_idx, mask=mask)
    tl.store(sorted_weights_ptr + write_pos, weights, mask=mask)
    tk_old = tl.atomic_add(token_counts_ptr + token_idx, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
    rev_pos = token_idx.to(tl.int64) * topk + tk_old.to(tl.int64)
    tl.store(reverse_idx_ptr + rev_pos, write_pos.to(tl.int32), mask=mask)


@triton.jit
def _fused_quant_sort_single_kernel(
    x_ptr, x_fp4_ptr, x_sc_ptr,
    stride_x_m, stride_x_n,
    stride_fp4_m, stride_fp4_n,
    stride_sc_m, stride_sc_n,
    topk_ids_ptr, topk_weights_ptr,
    sorted_token_idx_ptr, sorted_weights_ptr,
    counts_ptr, offsets_ptr, cum_blocks_ptr, write_counts_ptr,
    reverse_idx_ptr, token_counts_ptr,
    total_sorted, E, M, topk,
    scaleN, N_quant,
    QUANT_GRID: tl.constexpr,
    BM: tl.constexpr,
    BLOCK_E: tl.constexpr,
):
    pid = tl.program_id(0)
    if pid < QUANT_GRID:
        sxm = tl.cast(stride_x_m, tl.int64)
        sxn = tl.cast(stride_x_n, tl.int64)
        sfm = tl.cast(stride_fp4_m, tl.int64)
        sfn = tl.cast(stride_fp4_n, tl.int64)
        pid_m = pid // scaleN
        pid_n = pid % scaleN
        x_offs_m = pid_m * 128 + tl.arange(0, 128)
        x_offs_n = pid_n * 32 + tl.arange(0, 32)
        x_offs = x_offs_m[:, None] * sxm + x_offs_n[None, :] * sxn
        x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N_quant)[None, :]
        x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
        amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
        amax = amax.to(tl.int32, bitcast=True)
        amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        amax = amax.to(tl.float32, bitcast=True)
        scale_e8m0_unbiased = tl.log2(amax).floor() - 2
        scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
        quant_scale = tl.exp2(-scale_e8m0_unbiased)
        qx = x * quant_scale
        bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
        qx = qx.to(tl.uint32, bitcast=True)
        s = qx & 0x80000000
        e = (qx >> 23) & 0xFF
        m = qx & 0x7FFFFF
        E8_BIAS = 127
        E2_BIAS = 1
        adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
        m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
        e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
        e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
        e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
        e2m1_value = tl.reshape(e2m1_value, [128, 16, 2])
        evens, odds = tl.split(e2m1_value)
        out_tensor = evens | (odds << 4)
        out_offs_m = pid_m * 128 + tl.arange(0, 128)
        out_offs_n = pid_n * 16 + tl.arange(0, 16)
        out_offs = out_offs_m[:, None] * sfm + out_offs_n[None, :] * sfn
        out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N_quant // 2))[None, :]
        tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".cg")
        bs_offs_m = pid_m * 128 + tl.arange(0, 128)
        bs_offs_n = pid_n
        bs_offs = bs_offs_m[:, None] * stride_sc_m + bs_offs_n[None, :] * stride_sc_n
        bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N_quant)[None, :]
        tl.store(x_sc_ptr + bs_offs, bs_e8m0, mask=bs_mask, cache_modifier=".cg")
    else:
        idx = tl.arange(0, BLOCK_E)
        e_mask = idx < E
        tl.store(counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
        tl.store(write_counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
        e1_mask = idx < (E + 1)
        tl.store(offsets_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
        tl.store(cum_blocks_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
        offs = tl.arange(0, 256)
        t_mask = offs < M
        tl.store(token_counts_ptr + offs, tl.zeros([256], dtype=tl.int32), mask=t_mask)
        tl.debug_barrier()
        mask = offs < total_sorted
        ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
        weights = tl.load(topk_weights_ptr + offs, mask=mask, other=0.0)
        tl.atomic_add(counts_ptr + ids, tl.full([256], 1, dtype=tl.int32), mask=mask)
        tl.debug_barrier()
        counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
        token_offsets = tl.cumsum(counts, axis=0)
        tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_mask)
        n_blocks = tl.where(counts > 0, (counts + BM - 1) // BM, tl.zeros([BLOCK_E], dtype=tl.int64))
        block_offsets = tl.cumsum(n_blocks, axis=0)
        tl.store(cum_blocks_ptr + 1 + idx, block_offsets, mask=e_mask)
        tl.debug_barrier()
        token_idx = (offs // topk).to(tl.int32)
        old = tl.atomic_add(write_counts_ptr + ids, tl.full([256], 1, dtype=tl.int32), mask=mask)
        base = tl.load(offsets_ptr + ids, mask=mask, other=0).to(tl.int32)
        write_pos = (base + old).to(tl.int64)
        tl.store(sorted_token_idx_ptr + write_pos, token_idx, mask=mask)
        tl.store(sorted_weights_ptr + write_pos, weights, mask=mask)
        tk_old = tl.atomic_add(token_counts_ptr + token_idx, tl.full([256], 1, dtype=tl.int32), mask=mask)
        rev_pos = token_idx.to(tl.int64) * topk + tk_old.to(tl.int64)
        tl.store(reverse_idx_ptr + rev_pos, write_pos.to(tl.int32), mask=mask)


@triton.jit
def _fused_quant_count_kernel(
    x_ptr, x_fp4_ptr, x_sc_ptr,
    stride_x_m, stride_x_n,
    stride_fp4_m, stride_fp4_n,
    stride_sc_m, stride_sc_n,
    topk_ids_ptr, counts_ptr, offsets_ptr, cum_blocks_ptr, done_ptr,
    total_sorted, E, M, num_count_ctas,
    scaleN, N_quant,
    QUANT_GRID: tl.constexpr,
    BM: tl.constexpr,
    BLOCK_E: tl.constexpr,
):
    pid = tl.program_id(0)
    if pid < QUANT_GRID:
        sxm = tl.cast(stride_x_m, tl.int64)
        sxn = tl.cast(stride_x_n, tl.int64)
        sfm = tl.cast(stride_fp4_m, tl.int64)
        sfn = tl.cast(stride_fp4_n, tl.int64)
        pid_m = pid // scaleN
        pid_n = pid % scaleN
        x_offs_m = pid_m * 128 + tl.arange(0, 128)
        x_offs_n = pid_n * 32 + tl.arange(0, 32)
        x_offs = x_offs_m[:, None] * sxm + x_offs_n[None, :] * sxn
        x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N_quant)[None, :]
        x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
        amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
        amax = amax.to(tl.int32, bitcast=True)
        amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        amax = amax.to(tl.float32, bitcast=True)
        scale_e8m0_unbiased = tl.log2(amax).floor() - 2
        scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
        quant_scale = tl.exp2(-scale_e8m0_unbiased)
        qx = x * quant_scale
        bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
        qx = qx.to(tl.uint32, bitcast=True)
        s = qx & 0x80000000
        e = (qx >> 23) & 0xFF
        m = qx & 0x7FFFFF
        E8_BIAS = 127
        E2_BIAS = 1
        adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
        m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
        e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
        e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
        e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
        e2m1_value = tl.reshape(e2m1_value, [128, 16, 2])
        evens, odds = tl.split(e2m1_value)
        out_tensor = evens | (odds << 4)
        out_offs_m = pid_m * 128 + tl.arange(0, 128)
        out_offs_n = pid_n * 16 + tl.arange(0, 16)
        out_offs = out_offs_m[:, None] * sfm + out_offs_n[None, :] * sfn
        out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N_quant // 2))[None, :]
        tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".cg")
        bs_offs_m = pid_m * 128 + tl.arange(0, 128)
        bs_offs_n = pid_n
        bs_offs = bs_offs_m[:, None] * stride_sc_m + bs_offs_n[None, :] * stride_sc_n
        bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N_quant)[None, :]
        tl.store(x_sc_ptr + bs_offs, bs_e8m0, mask=bs_mask, cache_modifier=".cg")
    else:
        count_pid = pid - QUANT_GRID
        offs = count_pid * 256 + tl.arange(0, 256)
        mask = offs < total_sorted
        ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
        tl.atomic_add(counts_ptr + ids, tl.full([256], 1, dtype=tl.int32), mask=mask)
        old_done = tl.atomic_add(done_ptr, 1)
        if old_done == num_count_ctas - 1:
            idx = tl.arange(0, BLOCK_E)
            e_mask = idx < E
            counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
            token_offsets = tl.cumsum(counts, axis=0)
            tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_mask)
            n_blocks = tl.where(counts > 0, (counts + BM - 1) // BM, tl.zeros([BLOCK_E], dtype=tl.int64))
            block_offsets = tl.cumsum(n_blocks, axis=0)
            tl.store(cum_blocks_ptr + 1 + idx, block_offsets, mask=e_mask)


@triton.jit
def _scatter_reduce_kernel(
    src_ptr, reverse_idx_ptr, dst_ptr,
    M, dh,
    stride_sm, stride_sn,
    topk,
    BN: tl.constexpr, TOPK: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    stride_sm = tl.cast(stride_sm, tl.int64)
    stride_sn = tl.cast(stride_sn, tl.int64)
    cols = pid_n * BN + tl.arange(0, BN)
    col_mask = cols < dh
    acc = tl.zeros([BN], dtype=tl.float32)
    rev_base = pid_m * topk
    for k in range(TOPK):
        sorted_pos = tl.load(reverse_idx_ptr + rev_base + k).to(tl.int64)
        vals = tl.load(src_ptr + sorted_pos * stride_sm + cols * stride_sn, mask=col_mask, other=0.0).to(tl.float32)
        acc += vals
    tl.store(dst_ptr + pid_m.to(tl.int64) * dh + cols, acc.to(tl.bfloat16), mask=col_mask)


# ---- Batched MoE GEMM kernel (supports variable BM) ----

@triton.jit
def _batched_moe_gemm_fp4(
    a_ptr, b_ptr, c_ptr,
    a_sc_ptr, b_sc_ptr,
    token_idx_ptr, sorted_weights_ptr,
    cum_blocks_ptr, expert_offsets_ptr,
    E, N, K_half,
    stride_am, stride_ak,
    stride_be, stride_bn, stride_bk,
    stride_cm, stride_cn,
    stride_asm, stride_ask,
    stride_bse, stride_bsn, stride_bsk,
    max_m_blocks, total_tokens,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
    EVEN_K: tl.constexpr, INDIRECT: tl.constexpr,
    PRESHUFFLE: tl.constexpr, APPLY_WEIGHTS: tl.constexpr,
    SEARCH_ITERS: tl.constexpr, USE_CG: tl.constexpr = False,
    OUTPUT_BF16: tl.constexpr = False,
    NK: tl.constexpr = 0,
    EVEN_N: tl.constexpr = False,
):
    SG: tl.constexpr = 32
    stride_am = tl.cast(stride_am, tl.int64)
    stride_ak = tl.cast(stride_ak, tl.int64)
    stride_be = tl.cast(stride_be, tl.int64)
    stride_bn = tl.cast(stride_bn, tl.int64)
    stride_bk = tl.cast(stride_bk, tl.int64)
    stride_cm = tl.cast(stride_cm, tl.int64)
    stride_cn = tl.cast(stride_cn, tl.int64)
    stride_asm = tl.cast(stride_asm, tl.int64)
    stride_ask = tl.cast(stride_ask, tl.int64)
    stride_bse = tl.cast(stride_bse, tl.int64)
    stride_bsn = tl.cast(stride_bsn, tl.int64)
    stride_bsk = tl.cast(stride_bsk, tl.int64)
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid = tl.program_id(0)
    nn = tl.cdiv(N, BN)
    pid_mb = pid // nn
    pid_n = pid % nn
    if pid_mb >= max_m_blocks:
        return

    lo = tl.cast(0, tl.int64)
    hi = tl.cast(E, tl.int64)
    pid_mb_64 = tl.cast(pid_mb, tl.int64)
    for _ in range(SEARCH_ITERS):
        mid = (lo + hi + 1) // 2
        v = tl.load(cum_blocks_ptr + mid)
        cond = v <= pid_mb_64
        lo = tl.where(cond, mid, lo)
        hi = tl.where(cond, hi, mid - 1)
    expert_id = lo

    if expert_id >= E:
        return

    block_start = tl.load(cum_blocks_ptr + expert_id)
    local_block = pid_mb_64 - block_start
    expert_token_start = tl.load(expert_offsets_ptr + expert_id)
    expert_token_end = tl.load(expert_offsets_ptr + expert_id + 1)

    row_start = expert_token_start + local_block * BM
    if row_start >= expert_token_end:
        return

    sorted_rows = row_start + tl.arange(0, BM).to(tl.int64)
    row_mask = (sorted_rows < total_tokens) & (sorted_rows < expert_token_end)

    if INDIRECT:
        a_rows = tl.load(token_idx_ptr + sorted_rows, mask=row_mask, other=0).to(tl.int64)
    else:
        a_rows = sorted_rows

    cols = pid_n * BN + tl.arange(0, BN)
    col_mask = cols < N

    hk = tl.arange(0, BK // 2)
    ap = a_ptr + a_rows[:, None] * stride_am + hk[None, :] * stride_ak
    ks = tl.arange(0, BK // SG)
    asp = a_sc_ptr + a_rows[:, None] * stride_asm + ks[None, :] * stride_ask

    b_base = b_ptr + expert_id * stride_be
    bsc_base = b_sc_ptr + expert_id * stride_bse

    if PRESHUFFLE:
        offs_bn_shuf = pid_n * (BN // 16) + tl.arange(0, BN // 16)
        offs_k_shuf = tl.arange(0, (BK // 2) * 16)
        bp = b_base + offs_bn_shuf[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
        bsp = bsc_base + cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
    else:
        bp = b_base + hk[:, None] * stride_bk + cols[None, :].to(tl.int64) * stride_bn
        bsp = bsc_base + cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk

    acc = tl.zeros((BM, BN), dtype=tl.float32)

    if PRESHUFFLE and USE_CG:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            if EVEN_N:
                bsc = tl.load(bsp, cache_modifier=".cg")
            else:
                bsc = tl.load(bsp, mask=col_mask[:, None], other=0, cache_modifier=".cg")
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            b_raw = tl.load(bp, cache_modifier=".cg")
            b = (b_raw
                .reshape(1, BN // 16, BK // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BN, BK // 2)
                .trans(1, 0))
            acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp += (BK // 2) * 16 * stride_bk
            bsp += (BK // SG) * stride_bsk
    elif PRESHUFFLE:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            if EVEN_N:
                bsc = tl.load(bsp)
            else:
                bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            b_raw = tl.load(bp)
            b = (b_raw
                .reshape(1, BN // 16, BK // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BN, BK // 2)
                .trans(1, 0))
            acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp += (BK // 2) * 16 * stride_bk
            bsp += (BK // SG) * stride_bsk
    elif EVEN_K:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            b = tl.load(bp, mask=col_mask[None, :], other=0)
            acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
            ap += (BK // 2) * stride_ak
            bp += (BK // 2) * stride_bk
            asp += (BK // SG) * stride_ask
            bsp += (BK // SG) * stride_bsk
    else:
        K_rem = K_half
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None] & (hk[None, :] < K_rem), other=0)
            b = tl.load(bp, mask=(hk[:, None] < K_rem) & col_mask[None, :], other=0)
            acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
            ap += (BK // 2) * stride_ak
            bp += (BK // 2) * stride_bk
            asp += (BK // SG) * stride_ask
            bsp += (BK // SG) * stride_bsk
            K_rem -= BK // 2

    if APPLY_WEIGHTS:
        w = tl.load(sorted_weights_ptr + sorted_rows, mask=row_mask, other=0)
        acc = acc * w[:, None]

    cp = c_ptr + sorted_rows[:, None] * stride_cm + cols[None, :].to(tl.int64) * stride_cn
    store_mask = row_mask[:, None] if EVEN_N else (row_mask[:, None] & col_mask[None, :])
    if OUTPUT_BF16:
        tl.store(cp, acc.to(tl.bfloat16), mask=store_mask, cache_modifier=".cg")
    else:
        tl.store(cp, acc, mask=store_mask, cache_modifier=".cg")


@triton.jit
def _fused_gemm1_silu_quant_fp4(
    a_ptr, b_ptr,
    fp4_out_ptr, scale_out_ptr,
    a_sc_ptr, b_sc_ptr,
    token_idx_ptr,
    cum_blocks_ptr, expert_offsets_ptr,
    E, dep, K_half,
    stride_am, stride_ak,
    stride_be, stride_bn, stride_bk,
    stride_fp4m, stride_fp4n,
    stride_scm, stride_scn,
    stride_asm, stride_ask,
    stride_bse, stride_bsn, stride_bsk,
    max_m_blocks, total_tokens,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
    EVEN_K: tl.constexpr,
    SEARCH_ITERS: tl.constexpr, USE_CG: tl.constexpr = False,
    NK: tl.constexpr = 1,
    EVEN_N: tl.constexpr = False,
):
    SG: tl.constexpr = 32
    stride_am = tl.cast(stride_am, tl.int64)
    stride_ak = tl.cast(stride_ak, tl.int64)
    stride_be = tl.cast(stride_be, tl.int64)
    stride_bn = tl.cast(stride_bn, tl.int64)
    stride_bk = tl.cast(stride_bk, tl.int64)
    stride_fp4m = tl.cast(stride_fp4m, tl.int64)
    stride_fp4n = tl.cast(stride_fp4n, tl.int64)
    stride_asm = tl.cast(stride_asm, tl.int64)
    stride_ask = tl.cast(stride_ask, tl.int64)
    stride_bse = tl.cast(stride_bse, tl.int64)
    stride_bsn = tl.cast(stride_bsn, tl.int64)
    stride_bsk = tl.cast(stride_bsk, tl.int64)
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid = tl.program_id(0)
    nn = tl.cdiv(dep, BN)
    pid_mb = pid // nn
    pid_n = pid % nn
    if pid_mb >= max_m_blocks:
        return

    lo = tl.cast(0, tl.int64)
    hi = tl.cast(E, tl.int64)
    pid_mb_64 = tl.cast(pid_mb, tl.int64)
    for _ in range(SEARCH_ITERS):
        mid = (lo + hi + 1) // 2
        v = tl.load(cum_blocks_ptr + mid)
        cond = v <= pid_mb_64
        lo = tl.where(cond, mid, lo)
        hi = tl.where(cond, hi, mid - 1)
    expert_id = lo
    if expert_id >= E:
        return

    block_start = tl.load(cum_blocks_ptr + expert_id)
    local_block = pid_mb_64 - block_start
    expert_token_start = tl.load(expert_offsets_ptr + expert_id)
    expert_token_end = tl.load(expert_offsets_ptr + expert_id + 1)
    row_start = expert_token_start + local_block * BM
    if row_start >= expert_token_end:
        return

    sorted_rows = row_start + tl.arange(0, BM).to(tl.int64)
    row_mask = (sorted_rows < total_tokens) & (sorted_rows < expert_token_end)
    a_rows = tl.load(token_idx_ptr + sorted_rows, mask=row_mask, other=0).to(tl.int64)

    gate_cols = pid_n * BN + tl.arange(0, BN)
    up_cols = dep + gate_cols
    gate_col_mask = gate_cols < dep

    hk = tl.arange(0, BK // 2)
    ap = a_ptr + a_rows[:, None] * stride_am + hk[None, :] * stride_ak
    ks = tl.arange(0, BK // SG)
    asp = a_sc_ptr + a_rows[:, None] * stride_asm + ks[None, :] * stride_ask

    b_base = b_ptr + expert_id * stride_be
    bsc_base = b_sc_ptr + expert_id * stride_bse

    offs_bn_gate = pid_n * (BN // 16) + tl.arange(0, BN // 16)
    offs_bn_up = (dep // 16) + offs_bn_gate
    offs_k_shuf = tl.arange(0, (BK // 2) * 16)
    bp_gate = b_base + offs_bn_gate[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
    bp_up = b_base + offs_bn_up[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
    bsp_gate = bsc_base + gate_cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
    bsp_up = bsc_base + up_cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk

    acc_gate = tl.zeros((BM, BN), dtype=tl.float32)
    acc_up = tl.zeros((BM, BN), dtype=tl.float32)

    if USE_CG and EVEN_N:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            bsc_g = tl.load(bsp_gate, cache_modifier=".cg")
            b_raw_g = tl.load(bp_gate, cache_modifier=".cg")
            b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
            bsc_u = tl.load(bsp_up, cache_modifier=".cg")
            b_raw_u = tl.load(bp_up, cache_modifier=".cg")
            b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp_gate += (BK // 2) * 16 * stride_bk
            bp_up += (BK // 2) * 16 * stride_bk
            bsp_gate += (BK // SG) * stride_bsk
            bsp_up += (BK // SG) * stride_bsk
    elif USE_CG:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            bsc_g = tl.load(bsp_gate, mask=gate_col_mask[:, None], other=0, cache_modifier=".cg")
            b_raw_g = tl.load(bp_gate, cache_modifier=".cg")
            b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
            bsc_u = tl.load(bsp_up, mask=gate_col_mask[:, None], other=0, cache_modifier=".cg")
            b_raw_u = tl.load(bp_up, cache_modifier=".cg")
            b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp_gate += (BK // 2) * 16 * stride_bk
            bp_up += (BK // 2) * 16 * stride_bk
            bsp_gate += (BK // SG) * stride_bsk
            bsp_up += (BK // SG) * stride_bsk
    elif EVEN_N:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            bsc_g = tl.load(bsp_gate)
            b_raw_g = tl.load(bp_gate)
            b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
            bsc_u = tl.load(bsp_up)
            b_raw_u = tl.load(bp_up)
            b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp_gate += (BK // 2) * 16 * stride_bk
            bp_up += (BK // 2) * 16 * stride_bk
            bsp_gate += (BK // SG) * stride_bsk
            bsp_up += (BK // SG) * stride_bsk
    else:
        for _ in range(NK):
            asc = tl.load(asp, mask=row_mask[:, None], other=0)
            a = tl.load(ap, mask=row_mask[:, None], other=0)
            bsc_g = tl.load(bsp_gate, mask=gate_col_mask[:, None], other=0)
            b_raw_g = tl.load(bp_gate)
            b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
            bsc_u = tl.load(bsp_up, mask=gate_col_mask[:, None], other=0)
            b_raw_u = tl.load(bp_up)
            b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
            acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
            ap += (BK // 2) * stride_ak
            asp += (BK // SG) * stride_ask
            bp_gate += (BK // 2) * 16 * stride_bk
            bp_up += (BK // 2) * 16 * stride_bk
            bsp_gate += (BK // SG) * stride_bsk
            bsp_up += (BK // SG) * stride_bsk

    x = (acc_gate * tl.sigmoid(acc_gate) * acc_up).to(tl.bfloat16).to(tl.float32)
    amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_unb = tl.log2(amax).floor() - 2
    scale_unb = tl.clamp(scale_unb, min=-127, max=127)
    quant_scale = tl.exp2(-scale_unb)
    qx = x * quant_scale
    bs_e8m0 = scale_unb.to(tl.uint8) + 127
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    e = (qx >> 23) & 0xFF
    m = qx & 0x7FFFFF
    E8_BIAS: tl.constexpr = 127
    E2_BIAS: tl.constexpr = 1
    adj = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
    m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adj, m)
    e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
    e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
    e2m1_val = ((s >> 28) | e2m1_tmp).to(tl.uint8)
    e2m1_val = tl.reshape(e2m1_val, [BM, BN // 2, 2])
    evens, odds = tl.split(e2m1_val)
    packed = evens | (odds << 4)

    out_m = sorted_rows
    out_n = pid_n * BN // 2 + tl.arange(0, BN // 2)
    out_offs = out_m[:, None] * stride_fp4m + out_n[None, :] * stride_fp4n
    out_mask = row_mask[:, None] if EVEN_N else (row_mask[:, None] & (out_n < (dep // 2))[None, :])
    tl.store(fp4_out_ptr + out_offs, packed, mask=out_mask, cache_modifier=".cg")

    sc_n = pid_n
    sc_offs = out_m[:, None] * stride_scm + sc_n * stride_scn
    tl.store(scale_out_ptr + sc_offs, bs_e8m0, mask=row_mask[:, None], cache_modifier=".cg")


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

    dh = config["d_hidden"]
    dep = config["d_expert_pad"]
    dhp = config["d_hidden_pad"]
    M = hidden_states.shape[0]
    E = gate_up_weight.shape[0]
    topk = topk_ids.shape[1]
    dev = hidden_states.device
    total_sorted = M * topk

    if total_sorted == 0:
        return torch.zeros((M, dh), dtype=torch.bfloat16, device=dev)

    guw = gate_up_weight_shuffled.view(torch.uint8).view(E, (2 * dep) // 16, (dhp // 2) * 16)
    dw = down_weight_shuffled.view(torch.uint8).view(E, dhp // 16, (dep // 2) * 16)
    guw_sc = gate_up_weight_scale.view(torch.uint8).reshape(E, 2 * dep, dhp // 32)
    dw_sc = down_weight_scale.view(torch.uint8).reshape(E, dhp, dep // 32)

    guw_s0, guw_s1, guw_s2 = guw.stride()
    dw_s0, dw_s1, dw_s2 = dw.stride()
    guw_sc_s0, guw_sc_s1, guw_sc_s2 = guw_sc.stride()
    dw_sc_s0, dw_sc_s1, dw_sc_s2 = dw_sc.stride()

    hs = hidden_states

    estimated_per_expert = total_sorted / E
    if estimated_per_expert >= 64 and dep <= 512:
        BM_GEMM = 128
    elif estimated_per_expert >= 64:
        BM_GEMM = 64
    elif estimated_per_expert >= 8:
        BM_GEMM = 32
    else:
        BM_GEMM = 16

    flat_ids = topk_ids.reshape(-1)
    flat_weights = topk_weights.reshape(-1)

    scaleN_valid = triton.cdiv(dh, 32)
    scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
    pad_M_sc = triton.cdiv(M, 256) * 256
    i32_sz = 2 * E + M + 1
    i64_sz = 2 * (E + 1)
    scaleN_inter = triton.cdiv(dep, 32)
    use_fused = (BM_GEMM <= 16) or (BM_GEMM <= 32 and dep >= 512)

    A = 256
    def al(n):
        return (n + A - 1) & ~(A - 1)

    o = 0
    o_hfp4 = o;  s_hfp4 = M * (dh // 2);                 o += al(s_hfp4)
    o_hsc  = o;  s_hsc  = pad_M_sc * scaleN_pad;          o += al(s_hsc)
    o_sort = o
    o_si32 = o;  s_si32 = i32_sz * 4;                     o += al(s_si32)
    o_si64 = o;  s_si64 = i64_sz * 8;                     o += al(s_si64)
    o_sort_end = o
    o_stix = o;  s_stix = total_sorted * 4;               o += al(s_stix)
    o_sw   = o;  s_sw   = total_sorted * 2;               o += al(s_sw)
    o_rev  = o;  s_rev  = total_sorted * 4;               o += al(s_rev)
    o_ifp4 = o;  s_ifp4 = total_sorted * (dep // 2);      o += al(max(s_ifp4, 1))
    o_isc  = o;  s_isc  = total_sorted * scaleN_inter;    o += al(max(s_isc, 1))
    if not use_fused:
        o_g1 = o;  s_g1 = total_sorted * 2 * dep * 4;    o += al(s_g1)
    o_g2  = o;  s_g2  = total_sorted * dh * 2;            o += al(s_g2)
    o_out = o;  s_out = M * dh * 2;                       o += al(s_out)

    _buf = torch.empty(o, dtype=torch.uint8, device=dev)

    hs_fp4 = _buf[o_hfp4:o_hfp4 + s_hfp4].view(M, dh // 2)
    hs_sc_buf = _buf[o_hsc:o_hsc + s_hsc].view(pad_M_sc, scaleN_pad)
    sort_i32 = _buf[o_si32:o_si32 + s_si32].view(torch.int32)
    sort_i64 = _buf[o_si64:o_si64 + s_si64].view(torch.int64)
    sorted_token_idx = _buf[o_stix:o_stix + s_stix].view(torch.int32)
    sorted_weights = _buf[o_sw:o_sw + s_sw].view(torch.bfloat16)
    reverse_idx = _buf[o_rev:o_rev + s_rev].view(torch.int32)
    inter_fp4 = _buf[o_ifp4:o_ifp4 + s_ifp4].view(total_sorted, dep // 2) if s_ifp4 > 0 else _buf[o_ifp4:o_ifp4 + 1].view(1, 1)
    inter_sc = _buf[o_isc:o_isc + s_isc].view(total_sorted, scaleN_inter) if s_isc > 0 else _buf[o_isc:o_isc + 1].view(1, 1)
    if not use_fused:
        gemm1_out = _buf[o_g1:o_g1 + s_g1].view(torch.float32).view(total_sorted, 2 * dep)
    gemm2_out = _buf[o_g2:o_g2 + s_g2].view(torch.bfloat16).view(total_sorted, dh)
    out_bf16 = _buf[o_out:o_out + s_out].view(torch.bfloat16).view(M, dh)

    expert_counts = sort_i32[:E]
    write_counts = sort_i32[E:2*E]
    token_counts = sort_i32[2*E:2*E+M]
    done_counter = sort_i32[2*E+M:]
    expert_offsets = sort_i64[:E+1]
    cum_blocks = sort_i64[E+1:]

    SORT_BLOCK = 256
    BLOCK_E = triton.next_power_of_2(E)
    sort_grid = triton.cdiv(total_sorted, SORT_BLOCK)
    quant_grid_m = triton.cdiv(M, 128)
    quant_grid = quant_grid_m * scaleN_valid

    if sort_grid == 1:
        _fused_quant_sort_single_kernel[(quant_grid + 1,)](
            hs, hs_fp4, hs_sc_buf,
            *hs.stride(), *hs_fp4.stride(), *hs_sc_buf.stride(),
            flat_ids, flat_weights, sorted_token_idx, sorted_weights,
            expert_counts, expert_offsets, cum_blocks, write_counts,
            reverse_idx, token_counts,
            total_sorted, E, M, topk,
            scaleN_valid, dh,
            QUANT_GRID=quant_grid, BM=BM_GEMM, BLOCK_E=BLOCK_E,
            num_warps=4,
        )
    else:
        _buf[o_sort:o_sort_end].zero_()
        _fused_quant_count_kernel[(quant_grid + sort_grid,)](
            hs, hs_fp4, hs_sc_buf,
            *hs.stride(), *hs_fp4.stride(), *hs_sc_buf.stride(),
            flat_ids, expert_counts, expert_offsets, cum_blocks,
            done_counter, total_sorted, E, M, sort_grid,
            scaleN_valid, dh,
            QUANT_GRID=quant_grid, BM=BM_GEMM, BLOCK_E=BLOCK_E,
            num_warps=4,
        )
        _moe_scatter_kernel[(sort_grid,)](
            flat_ids, flat_weights,
            sorted_token_idx, sorted_weights,
            expert_offsets, write_counts,
            reverse_idx, token_counts,
            topk, total_sorted, BLOCK=SORT_BLOCK,
        )
    hs_sc = hs_sc_buf[:M]

    max_m_blocks = min(total_sorted, (total_sorted + BM_GEMM - 1) // BM_GEMM + E)
    si = 6 if E <= 64 else 9
    nw1 = 4 if BM_GEMM >= 64 else 2
    max_bk = 512 if BM_GEMM >= 64 else 1024
    use_mfma16 = BM_GEMM <= 32

    BK1 = 128
    for bk in [1024, 512, 256, 128]:
        if bk <= min(dh, max_bk) and dh % bk == 0:
            BK1 = bk
            break
    even_k1 = (dh % BK1) == 0
    NK1 = triton.cdiv(dh // 2, BK1 // 2)

    if use_fused:
        BN_fused = 32
        nn_fused = triton.cdiv(dep, BN_fused)
        _fused_gemm1_silu_quant_fp4[(max_m_blocks * nn_fused,)](
            hs_fp4, guw,
            inter_fp4, inter_sc,
            hs_sc, guw_sc,
            sorted_token_idx,
            cum_blocks, expert_offsets,
            E, dep, dh // 2,
            hs_fp4.stride(0), hs_fp4.stride(1),
            guw_s0, guw_s1, guw_s2,
            inter_fp4.stride(0), inter_fp4.stride(1),
            inter_sc.stride(0), inter_sc.stride(1),
            hs_sc.stride(0), hs_sc.stride(1),
            guw_sc_s0, guw_sc_s1, guw_sc_s2,
            max_m_blocks, total_sorted,
            BM=BM_GEMM, BN=BN_fused, BK=BK1,
            EVEN_K=even_k1,
            SEARCH_ITERS=si, USE_CG=(dep <= 512),
            NK=NK1,
            EVEN_N=(dep % BN_fused == 0),
            num_warps=nw1, num_stages=(1 if BM_GEMM >= 32 else 2),
            matrix_instr_nonkdim=16, schedule_hint="attention",
        )
    else:
        if dep > 512:
            BN1 = 256
        elif dep <= 256 and BM_GEMM <= 32 and total_sorted < 2000:
            BN1 = 128
        elif dep == 512 and BM_GEMM >= 64:
            BN1 = 128
        else:
            BN1 = 64

        nn1 = triton.cdiv(2 * dep, BN1)
        _batched_moe_gemm_fp4[(max_m_blocks * nn1,)](
            hs_fp4, guw, gemm1_out,
            hs_sc, guw_sc,
            sorted_token_idx, sorted_weights,
            cum_blocks, expert_offsets,
            E, 2 * dep, dh // 2,
            hs_fp4.stride(0), hs_fp4.stride(1),
            guw_s0, guw_s1, guw_s2,
            gemm1_out.stride(0), gemm1_out.stride(1),
            hs_sc.stride(0), hs_sc.stride(1),
            guw_sc_s0, guw_sc_s1, guw_sc_s2,
            max_m_blocks, total_sorted,
            BM=BM_GEMM, BN=BN1, BK=BK1, EVEN_K=even_k1, INDIRECT=True,
            PRESHUFFLE=True, APPLY_WEIGHTS=False,
            SEARCH_ITERS=si, USE_CG=(dep <= 512),
            OUTPUT_BF16=False, NK=NK1,
            EVEN_N=((2 * dep) % BN1 == 0),
            num_warps=(2 if BM_GEMM >= 128 else nw1), num_stages=(3 if dep == 512 and BM_GEMM >= 64 else 2),
            **({"matrix_instr_nonkdim": 16} if use_mfma16 else {}),
            schedule_hint="attention",
        )
        _fused_silu_mul_quant_kernel[(triton.cdiv(total_sorted, 128), scaleN_inter)](
            gemm1_out, inter_fp4, inter_sc,
            dep,
            gemm1_out.stride(0), gemm1_out.stride(1),
            inter_fp4.stride(0), inter_fp4.stride(1),
            inter_sc.stride(0), inter_sc.stride(1),
            total_sorted,
            BLOCK_M=128, QB=32,
            num_stages=3,
        )

    BN2 = 256 if dep <= 256 else 128
    nw2 = 4
    BK2 = 128
    bk2_max = 256 if dep == 512 else 512
    for bk in [512, 256, 128]:
        if bk <= min(dep, bk2_max) and dep % bk == 0:
            BK2 = bk
            break
    even_k2 = (dep % BK2) == 0
    NK2 = triton.cdiv(dep // 2, BK2 // 2)

    nn2 = triton.cdiv(dh, BN2)
    _batched_moe_gemm_fp4[(max_m_blocks * nn2,)](
        inter_fp4, dw, gemm2_out,
        inter_sc, dw_sc,
        sorted_token_idx, sorted_weights,
        cum_blocks, expert_offsets,
        E, dh, dep // 2,
        inter_fp4.stride(0), inter_fp4.stride(1),
        dw_s0, dw_s1, dw_s2,
        gemm2_out.stride(0), gemm2_out.stride(1),
        inter_sc.stride(0), inter_sc.stride(1),
        dw_sc_s0, dw_sc_s1, dw_sc_s2,
        max_m_blocks, total_sorted,
        BM=BM_GEMM, BN=BN2, BK=BK2, EVEN_K=even_k2, INDIRECT=False,
        PRESHUFFLE=True, APPLY_WEIGHTS=True,
        SEARCH_ITERS=si, USE_CG=(dep <= 512),
        OUTPUT_BF16=True, NK=NK2,
        EVEN_N=(dh % BN2 == 0),
        num_warps=nw2, num_stages=2,
        **({"matrix_instr_nonkdim": 16} if use_mfma16 else {}),
        schedule_hint="attention",
    )

    SR_BN = 256 if M >= 64 else 128
    _scatter_reduce_kernel[(M, triton.cdiv(dh, SR_BN))](
        gemm2_out, reverse_idx, out_bf16,
        M, dh,
        gemm2_out.stride(0), gemm2_out.stride(1),
        topk,
        BN=SR_BN, TOPK=topk,
        num_warps=4, num_stages=3,
    )
    return out_bf16
scrolls · 1141 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