Skip to content
KernelIndex
Search⌘K

submission 696235

dc1312 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-696235?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
207.0µs
#767 of 782
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d704331211ecf01987023abce2c5684c5d846b1b5805b2ddc8d9962efdfdd3ff
license declaredunknown
license concludedunknown
authorsdc1312
imported2026-08-26

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-epilogue- Inline fused_silu_mul_quant with cached buffers
num-warps = 1num_warps=1,
stages = 2num_warps=nw1, num_stages=2,
tile-m = 128BLOCK_M=128, QB=32,
tile-n = 128SR_BN = 128

Kernel source

submission.py603 lines
"""
v67: Grid-aware BN1 + inlined allocs + combined v60 optimizations.
- BN1=128 for dep=256 when total_sorted<2000 (small grid → need big CTAs)
- BN1=64 for dep=256 when total_sorted>=2000 (enough parallelism → more N-tiles)
- Inline fused_silu_mul_quant with cached buffers
- All v60 sort/scatter/buffer optimizations
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl

_BUF = {}

def _get_buf(key, numel, dtype, device):
    if key not in _BUF or _BUF[key].numel() < numel or _BUF[key].dtype != dtype:
        _BUF[key] = torch.empty(numel, dtype=dtype, device=device)
    return _BUF[key][:numel]


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


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


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_kernel(topk_ids_ptr, counts_ptr, 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)
    tl.atomic_add(counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)


@triton.jit
def _moe_offsets_kernel(
    counts_ptr, offsets_ptr, cum_blocks_ptr,
    E,
    BM: tl.constexpr,
    BLOCK_E: tl.constexpr,
):
    idx = tl.arange(0, BLOCK_E)
    mask = idx < E
    counts = tl.load(counts_ptr + idx, mask=mask, other=0).to(tl.int64)
    token_offsets = tl.cumsum(counts, axis=0)
    tl.store(offsets_ptr + 1 + idx, token_offsets, mask=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=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 _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)
        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,
):
    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_raw = tl.program_id(0)
    nn = tl.cdiv(N, BN)
    grid_mn = max_m_blocks * nn
    pid = _remap_xcd(pid_raw, grid_mn)
    pid_n = pid // max_m_blocks
    pid_mb = pid % max_m_blocks

    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)
    nk = tl.cdiv(K_half, BK // 2)

    if PRESHUFFLE:
        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_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
    tl.store(cp, acc, mask=row_mask[:, None] & col_mask[None, :])


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)

    hs = F.pad(hidden_states, (0, dhp - dh)) if dhp > dh else hidden_states
    hs_fp4, hs_sc = dynamic_mxfp4_quant(hs)

    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 >= 4:
        BM_GEMM = 32
    else:
        BM_GEMM = 16

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

    # Consolidated int32 sort buffers: [expert_counts(E) | write_counts(E) | token_counts(M)]
    i32_sz = 2 * E + M
    sort_i32 = _get_buf("si32", i32_sz, torch.int32, dev)
    sort_i32.zero_()
    expert_counts = sort_i32[:E]
    write_counts = sort_i32[E:2*E]
    token_counts = sort_i32[2*E:]

    # Consolidated int64 sort buffers: [expert_offsets(E+1) | cum_blocks(E+1)]
    i64_sz = 2 * (E + 1)
    sort_i64 = _get_buf("si64", i64_sz, torch.int64, dev)
    sort_i64.zero_()
    expert_offsets = sort_i64[:E+1]
    cum_blocks = sort_i64[E+1:]

    sorted_token_idx = _get_buf("sti", total_sorted, torch.int32, dev)
    sorted_weights = _get_buf("sw", total_sorted, flat_weights.dtype, dev)
    reverse_idx = _get_buf("ri", M * topk, torch.int32, dev)

    SORT_BLOCK = 256
    _moe_count_kernel[(triton.cdiv(total_sorted, SORT_BLOCK),)](
        flat_ids, expert_counts, total_sorted, BLOCK=SORT_BLOCK,
    )

    BLOCK_E = triton.next_power_of_2(E)
    _moe_offsets_kernel[(1,)](
        expert_counts, expert_offsets, cum_blocks,
        E, BM=BM_GEMM, BLOCK_E=BLOCK_E,
        num_warps=1,
    )

    _moe_scatter_kernel[(triton.cdiv(total_sorted, SORT_BLOCK),)](
        flat_ids, flat_weights,
        sorted_token_idx, sorted_weights,
        expert_offsets, write_counts,
        reverse_idx, token_counts,
        topk, total_sorted, BLOCK=SORT_BLOCK,
    )

    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

    # BN1=128 for dep=256 when grid is small; BN1=64 when enough parallelism
    BN1 = 128 if (dep <= 256 and BM_GEMM <= 32 and total_sorted < 2000) else 64
    BK1 = 128
    for bk in [1024, 512, 256, 128]:
        if bk <= min(dhp, max_bk) and dhp % bk == 0:
            BK1 = bk
            break
    even_k1 = (dhp % BK1) == 0

    f32_sz = total_sorted * max(2 * dep, dhp)
    gemm_f32 = _get_buf("gf32", f32_sz, torch.float32, dev)
    gemm1_out = gemm_f32[:total_sorted * 2 * dep].view(total_sorted, 2 * dep)
    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, dhp // 2,
        hs_fp4.stride(0), hs_fp4.stride(1),
        guw.stride(0), guw.stride(1), guw.stride(2),
        gemm1_out.stride(0), gemm1_out.stride(1),
        hs_sc.stride(0), hs_sc.stride(1),
        guw_sc.stride(0), guw_sc.stride(1), guw_sc.stride(2),
        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,
        num_warps=nw1, num_stages=2,
    )

    # Inline fused_silu_mul_quant with cached buffers
    inter_fp4 = _get_buf("ifp4", total_sorted * dep // 2, torch.uint8, dev).view(total_sorted, dep // 2)
    scaleN_inter = triton.cdiv(dep, 32)
    inter_sc = _get_buf("isc", total_sorted * scaleN_inter, torch.uint8, dev).view(total_sorted, scaleN_inter)
    _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,
    )

    # GEMM2: BN2=256 for dep=256 (BK=256, fits LDS), BN2=128 otherwise
    BN2 = 256 if dep <= 256 else 128
    nw2 = 4
    BK2 = 128
    for bk in [512, 256, 128]:
        if bk <= dep and dep % bk == 0:
            BK2 = bk
            break
    even_k2 = (dep % BK2) == 0

    gemm2_out = gemm_f32[:total_sorted * dhp].view(total_sorted, dhp)
    nn2 = triton.cdiv(dhp, 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, dhp, dep // 2,
        inter_fp4.stride(0), inter_fp4.stride(1),
        dw.stride(0), dw.stride(1), dw.stride(2),
        gemm2_out.stride(0), gemm2_out.stride(1),
        inter_sc.stride(0), inter_sc.stride(1),
        dw_sc.stride(0), dw_sc.stride(1), dw_sc.stride(2),
        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,
        num_warps=nw2, num_stages=2,
    )

    out_bf16 = _get_buf("out_bf16", M * dh, torch.bfloat16, dev).view(M, dh)
    SR_BN = 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,
    )
    return out_bf16
scrolls · 603 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