Skip to content
KernelIndex
Search⌘K

submission 754587

ooousay · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754587?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.71µs
#82 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5ee9f51ffa298859b904a97ff234a1dabb7831a48ad90e943bf5534a9fd197fa
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4, num_stages=2,
split-k"""M=16, N=2112, K=7168 ? custom hardcoded split-K kernel with native hw quant."""
stages = 2num_warps=4, num_stages=2,
tile-k = 256BLOCK_SIZE_K=256,
tile-m = 4BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,
tile-n = 128BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,

Kernel source

kernel.py1020 lines
#!POPCORN leaderboard amd-mxfp4-mm
"""Auto-generated by build.py"""

# ============================================================
# m=4, n=2880, k=512
# ============================================================

"""M=4, N=2880, K=512 ? constexpr shapes (no binary patching)."""
import torch
import triton
import triton.language as tl


@triton.jit
def _mxfp4_quant_op_native(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """Native hw quant using v_cvt_scalef32_pk_fp4_bf16 -- shared by all shapes."""
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Scale computation
    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)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    # Native convert+round+pack via hardware instruction
    hw_scale = tl.exp2(scale_e8m0_unbiased)
    x_bf16 = x.to(tl.bfloat16)
    x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    x_lo, x_hi = tl.split(x_pairs)
    x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
    x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
    packed_bf16x2 = x_lo_i32 | x_hi_i32

    hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))

    hw_result = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v,v,v",
        [packed_bf16x2, hw_scale_bc],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )
    x_fp4 = (hw_result & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _hardcoded_m4_kernel_4_2880_512(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_cm, stride_cn, stride_bsn, stride_bsk,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    M: tl.constexpr = 4
    N: tl.constexpr = 2880
    K: tl.constexpr = 256       # k // 2
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_PID_N: tl.constexpr = 23  # cdiv(2880, 128)
    NUM_K_ITER: tl.constexpr = 2  # cdiv(512 // 2, 256 // 2)

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid = tl.program_id(axis=0)
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N
    tl.assume(pid_m >= 0)
    tl.assume(pid_n >= 0)

    # A offsets (constant across k)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)

    # B offsets (constant across k)
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)

    # B scale offsets (constant across k)
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)

    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    for k in range(NUM_K_ITER):
        # Recompute pointers each iteration from function args
        # (enables ConvertToBufferOps ? buffer_load)
        a_offs = (
            offs_am[:, None] * stride_am
            + (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
        )
        b_offs = (
            offs_bn[:, None] * stride_bn
            + (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
        )
        bs_offs = (
            offs_bsn[:, None] * stride_bsn
            + (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
        )

        b_scales = (
            tl.load(b_scales_ptr + bs_offs, cache_modifier=cache_modifier)
            .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        )

        a_bf16 = tl.load(a_ptr + a_offs)
        b = tl.load(b_ptr + b_offs, cache_modifier=cache_modifier)

        b = (
            b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )

        a, a_scales = _mxfp4_quant_op_native(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
        accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

    c = accumulator.to(c_ptr.type.element_ty)
    offs_cm = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
    offs_cn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)).to(tl.int64)
    c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    tl.store(c_ptrs, c, mask=c_mask)


_state_4_2880_512 = None

def _run_4_2880_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_4_2880_512
    if _state_4_2880_512 is None:
        _state_4_2880_512 = {
            'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
            'B_w': None, 'B_sc': None, '_b_ptr': None,
        }

    b_ptr = B_shuffle.data_ptr()
    if _state_4_2880_512['_b_ptr'] != b_ptr:
        _state_4_2880_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
        _state_4_2880_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
        _state_4_2880_512['_b_ptr'] = b_ptr

    _hardcoded_m4_kernel_4_2880_512[(23,)](
        A, _state_4_2880_512['B_w'], _state_4_2880_512['out'], _state_4_2880_512['B_sc'],
        A.stride(0), A.stride(1),
        _state_4_2880_512['B_w'].stride(0), _state_4_2880_512['B_w'].stride(1),
        _state_4_2880_512['out'].stride(0), _state_4_2880_512['out'].stride(1),
        _state_4_2880_512['B_sc'].stride(0), _state_4_2880_512['B_sc'].stride(1),
        BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,
        BLOCK_SIZE_K=256,
        num_warps=4, num_stages=2,
        waves_per_eu=0, matrix_instr_nonkdim=16,
        cache_modifier=".cg",
    )

    return _state_4_2880_512['out']


# ============================================================
# m=16, n=2112, k=7168
# ============================================================

"""M=16, N=2112, K=7168 ? custom hardcoded split-K kernel with native hw quant."""
import torch
import triton
import triton.language as tl
from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk


@triton.jit
def _mxfp4_quant_op_native_16_2112_7168(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Scale computation
    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)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    # Native convert+round+pack via hardware instruction
    hw_scale = tl.exp2(scale_e8m0_unbiased)
    x_bf16 = x.to(tl.bfloat16)
    x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    x_lo, x_hi = tl.split(x_pairs)
    x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
    x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
    packed_bf16x2 = x_lo_i32 | x_hi_i32

    hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))

    hw_result = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v,v,v",
        [packed_bf16x2, hw_scale_bc],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )
    x_fp4 = (hw_result & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _hardcoded_preshuffle_splitk_kernel_16_2112_7168(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
    NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    # Hardcoded shape constants
    M: tl.constexpr = 16
    N: tl.constexpr = 2112
    K: tl.constexpr = 3584       # k // 2
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_PID_N: tl.constexpr = 17  # cdiv(2112, 128)
    NUM_K_ITER: tl.constexpr = 2  # cdiv(SPLITK_BLOCK_SIZE//2, BLOCK_SIZE_K//2)

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_ck > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid_unified = tl.program_id(axis=0)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N

    tl.assume(pid_m >= 0)
    tl.assume(pid_n >= 0)
    tl.assume(pid_k >= 0)

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        # A offsets (constant across k)
        offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)

        # B offsets (constant across k)
        offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)

        # B scale offsets (constant across k)
        offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
        offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)

        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for k in range(pid_k * NUM_K_ITER, (pid_k + 1) * NUM_K_ITER):
            # Recompute pointers each iteration from function args
            # (enables ConvertToBufferOps)
            a_offs = (
                offs_am[:, None] * stride_am
                + (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
            )
            b_offs = (
                offs_bn[:, None] * stride_bn
                + (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
            )
            bs_offs = (
                offs_bsn[:, None] * stride_bsn
                + (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
            )

            b_scales = (
                tl.load(b_scales_ptr + bs_offs, cache_modifier=cache_modifier)
                .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )

            a_bf16 = tl.load(a_ptr + a_offs)
            b = tl.load(b_ptr + b_offs, cache_modifier=cache_modifier)

            b = (
                b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )

            a, a_scales = _mxfp4_quant_op_native_16_2112_7168(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

        c = accumulator.to(c_ptr.type.element_ty)
        offs_cm = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
        offs_cn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)).to(tl.int64)
        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        tl.store(c_ptrs, c, mask=c_mask)


_state_16_2112_7168 = None

def _run_16_2112_7168(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_16_2112_7168
    if _state_16_2112_7168 is None:
        K_kernel = k // 2
        SPLITK_BLOCK_SIZE, BSK, actual_ksplit = get_splitk(K_kernel, 256, 14)
        grid_size = actual_ksplit * triton.cdiv(m, 16) * triton.cdiv(n, 128)
        y_pp = torch.empty((actual_ksplit, m, n), dtype=torch.float32, device=A.device)

        _state_16_2112_7168 = {
            'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
            'B_w': None, 'B_sc': None, '_b_ptr': None,
            'y_pp': y_pp,
            'grid_size': grid_size,
            'K_kernel': K_kernel,
            'BSK': BSK,
            'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE,
            'NUM_KSPLIT': actual_ksplit,
            'ACTUAL_KSPLIT': triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2)),
            'MAX_KSPLIT': triton.next_power_of_2(actual_ksplit),
            'reduce_grid': (triton.cdiv(m, 16), triton.cdiv(n, 64)),
        }

    s = _state_16_2112_7168

    b_ptr = B_shuffle.data_ptr()
    if s['_b_ptr'] != b_ptr:
        s['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
        s['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
        s['_b_ptr'] = b_ptr

    _hardcoded_preshuffle_splitk_kernel_16_2112_7168[(s['grid_size'],)](
        A, s['B_w'], s['y_pp'], s['B_sc'],
        A.stride(0), A.stride(1),
        s['B_w'].stride(0), s['B_w'].stride(1),
        s['y_pp'].stride(0), s['y_pp'].stride(1), s['y_pp'].stride(2),
        s['B_sc'].stride(0), s['B_sc'].stride(1),
        BLOCK_SIZE_M=16, BLOCK_SIZE_N=128,
        BLOCK_SIZE_K=s['BSK'],
        NUM_KSPLIT=s['NUM_KSPLIT'], SPLITK_BLOCK_SIZE=s['SPLITK_BLOCK_SIZE'],
        num_warps=4, num_stages=2,
        waves_per_eu=2, matrix_instr_nonkdim=16,
        cache_modifier=".cg",
    )

    _gluon_reduce_kernel[s['reduce_grid']](
        s['y_pp'], s['out'],
        m, n,
        s['y_pp'].stride(0), s['y_pp'].stride(1), s['y_pp'].stride(2),
        s['out'].stride(0), s['out'].stride(1),
        16, 64,
        s['ACTUAL_KSPLIT'], s['MAX_KSPLIT'],
    )

    return s['out']


# ============================================================
# m=32, n=4096, k=512
# ============================================================

"""M=32, N=4096, K=512 ? v13: custom hardcoded preshuffle, all constexpr."""
import torch
import triton
import triton.language as tl
from aiter import dtypes


@triton.jit
def _hardcoded_preshuffle_kernel_32_4096_512(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_cm,
    stride_cn,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    # Shape constants ? M=32, N=4096, K=512
    M: tl.constexpr = 32
    N: tl.constexpr = 4096
    K: tl.constexpr = 256  # K_kernel = 512 // 2
    SCALE_GROUP_SIZE: tl.constexpr = 32

    # Grid: 4 x 32 = 128 WGs
    NUM_PID_M: tl.constexpr = 4   # cdiv(32, 8)
    NUM_PID_N: tl.constexpr = 32  # cdiv(4096, 128)
    NUM_K_ITER: tl.constexpr = 2  # (512//2) / (256//2)

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid = tl.program_id(axis=0)
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N

    tl.assume(pid_m >= 0)
    tl.assume(pid_n >= 0)

    # A pointers ? no mask (32%8=0)
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    a_ptrs = a_ptr + (
        offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
    )

    # B pointers (preshuffled) ? no mask (4096%128=0)
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    b_ptrs = b_ptr + (
        offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
    )

    # B scale pointers
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
    b_scale_ptrs = (
        b_scales_ptr
        + offs_bsn[:, None] * stride_bsn
        + offs_ks[None, :] * stride_bsk
    )

    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    for k in range(NUM_K_ITER):
        b_scales = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .reshape(
                BLOCK_SIZE_N // 32,
                BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                4,
                16,
                2,
                2,
                1,
            )
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        )

        a_bf16 = tl.load(a_ptrs)
        b = tl.load(b_ptrs, cache_modifier=cache_modifier)

        b = (
            b.reshape(
                1,
                BLOCK_SIZE_N // 16,
                BLOCK_SIZE_K // 64,
                2,
                16,
                16,
            )
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )

        a, a_scales = _mxfp4_quant_op_native(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
        accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

        a_ptrs += BLOCK_SIZE_K * stride_ak
        b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
        b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

    c = accumulator.to(c_ptr.type.element_ty)

    offs_cm = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
    offs_cn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)).to(tl.int64)
    c_ptrs = (
        c_ptr
        + stride_cm * offs_cm[:, None]
        + stride_cn * offs_cn[None, :]
    )
    tl.store(c_ptrs, c)


_state_32_4096_512 = None

def _run_32_4096_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_32_4096_512
    if _state_32_4096_512 is None:
        _state_32_4096_512 = {
            'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
            'B_w': None, 'B_sc': None, '_b_ptr': None,
        }

    b_ptr = B_shuffle.data_ptr()
    if _state_32_4096_512['_b_ptr'] != b_ptr:
        _state_32_4096_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
        _state_32_4096_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
        _state_32_4096_512['_b_ptr'] = b_ptr

    _hardcoded_preshuffle_kernel_32_4096_512[(128,)](
        A, _state_32_4096_512['B_w'], _state_32_4096_512['out'], _state_32_4096_512['B_sc'],
        A.stride(0), A.stride(1),
        _state_32_4096_512['B_w'].stride(0), _state_32_4096_512['B_w'].stride(1),
        _state_32_4096_512['out'].stride(0), _state_32_4096_512['out'].stride(1),
        _state_32_4096_512['B_sc'].stride(0), _state_32_4096_512['B_sc'].stride(1),
        BLOCK_SIZE_M=8, BLOCK_SIZE_N=128,
        BLOCK_SIZE_K=256,
        num_warps=4, num_stages=2,
        waves_per_eu=2, matrix_instr_nonkdim=16,
        cache_modifier=".cg",
    )
    return _state_32_4096_512['out']


# ============================================================
# m=32, n=2880, k=512
# ============================================================

"""M=32, N=2880, K=512 ? fused_direct path. BSM=8 BSN=128 BSK=256."""
import torch
import triton
from aiter import dtypes
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
)

_state_32_2880_512 = None

def _run_32_2880_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_32_2880_512
    if _state_32_2880_512 is None:
        _state_32_2880_512 = {
            'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
            'B_w': None, 'B_sc': None, '_b_ptr': None,
            'grid_size': triton.cdiv(m, 8) * triton.cdiv(n, 128),
            'K_kernel': k // 2,
            'SPLITK_BLOCK_SIZE': k,
        }

    b_ptr = B_shuffle.data_ptr()
    if _state_32_2880_512['_b_ptr'] != b_ptr:
        _state_32_2880_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
        _state_32_2880_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
        _state_32_2880_512['_b_ptr'] = b_ptr

    _gemm_a16wfp4_preshuffle_kernel[(_state_32_2880_512['grid_size'],)](
        A, _state_32_2880_512['B_w'], _state_32_2880_512['out'], _state_32_2880_512['B_sc'],
        m, n, _state_32_2880_512['K_kernel'],
        A.stride(0), A.stride(1),
        _state_32_2880_512['B_w'].stride(0), _state_32_2880_512['B_w'].stride(1),
        0, _state_32_2880_512['out'].stride(0), _state_32_2880_512['out'].stride(1),
        _state_32_2880_512['B_sc'].stride(0), _state_32_2880_512['B_sc'].stride(1),
        BLOCK_SIZE_M=8, BLOCK_SIZE_N=128,
        BLOCK_SIZE_K=256, GROUP_SIZE_M=1,
        NUM_KSPLIT=1, SPLITK_BLOCK_SIZE=_state_32_2880_512['SPLITK_BLOCK_SIZE'],
        num_warps=4, num_stages=2,
        waves_per_eu=2, matrix_instr_nonkdim=16,
        PREQUANT=True, cache_modifier=None,
    )
    return _state_32_2880_512['out']


# ============================================================
# m=64, n=7168, k=2048
# ============================================================

"""M=64, N=7168, K=2048 ? v14 with native v_cvt_scalef32_pk_fp4_bf16 quant."""
import torch
import triton
import triton.language as tl
from aiter import dtypes


@triton.jit
def _mxfp4_quant_op_native_64_7168_2048(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Scale computation ? identical to aiter's _mxfp4_quant_op
    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)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    # Native convert+round+pack
    hw_scale = tl.exp2(scale_e8m0_unbiased)
    x_bf16 = x.to(tl.bfloat16)
    x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    x_lo, x_hi = tl.split(x_pairs)
    x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
    x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
    packed_bf16x2 = x_lo_i32 | x_hi_i32

    hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))

    hw_result = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v,v,v",
        [packed_bf16x2, hw_scale_bc],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )
    x_fp4 = (hw_result & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _hardcoded_preshuffle_kernel_64_7168_2048(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_cm,
    stride_cn,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    M: tl.constexpr = 64
    N: tl.constexpr = 7168
    K: tl.constexpr = 1024
    SCALE_GROUP_SIZE: tl.constexpr = 32
    SPLITK_BLOCK_SIZE: tl.constexpr = 2048

    NUM_PID_M: tl.constexpr = 4
    NUM_PID_N: tl.constexpr = 56
    NUM_K_ITER: tl.constexpr = 8

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid = tl.program_id(axis=0)
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N

    tl.assume(pid_m >= 0)
    tl.assume(pid_n >= 0)

    # A: row offsets (constant across k)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)

    # B: n-group offsets (constant across k)
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)

    # B scales: n-group offsets (constant across k)
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)

    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    for k in range(NUM_K_ITER):
        # Recompute pointers each iteration: splat(func_arg) + tensor_offset
        # Enables ConvertToBufferOps to match and emit buffer_load
        a_offs = (
            offs_am[:, None] * stride_am
            + (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
        )
        b_offs = (
            offs_bn[:, None] * stride_bn
            + (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
        )
        bs_offs = (
            offs_bsn[:, None] * stride_bsn
            + (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
        )

        b_scales = (
            tl.load(b_scales_ptr + bs_offs, cache_modifier=cache_modifier)
            .reshape(
                BLOCK_SIZE_N // 32,
                BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                4,
                16,
                2,
                2,
                1,
            )
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        )

        a_bf16 = tl.load(a_ptr + a_offs)
        b = tl.load(b_ptr + b_offs, cache_modifier=cache_modifier)

        b = (
            b.reshape(
                1,
                BLOCK_SIZE_N // 16,
                BLOCK_SIZE_K // 64,
                2,
                16,
                16,
            )
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )

        a, a_scales = _mxfp4_quant_op_native_64_7168_2048(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

        accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

    c = accumulator.to(c_ptr.type.element_ty)

    offs_cm = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
    offs_cn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)).to(tl.int64)
    c_ptrs = (
        c_ptr
        + stride_cm * offs_cm[:, None]
        + stride_cn * offs_cn[None, :]
    )
    tl.store(c_ptrs, c)


_state_64_7168_2048 = None

def _run_64_7168_2048(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_64_7168_2048
    if _state_64_7168_2048 is None:
        _state_64_7168_2048 = {
            'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
            'B_w': None, 'B_sc': None, '_b_ptr': None,
        }

    b_ptr = B_shuffle.data_ptr()
    if _state_64_7168_2048['_b_ptr'] != b_ptr:
        _state_64_7168_2048['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
        _state_64_7168_2048['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
        _state_64_7168_2048['_b_ptr'] = b_ptr

    _hardcoded_preshuffle_kernel_64_7168_2048[(224,)](
        A, _state_64_7168_2048['B_w'], _state_64_7168_2048['out'], _state_64_7168_2048['B_sc'],
        A.stride(0), A.stride(1),
        _state_64_7168_2048['B_w'].stride(0), _state_64_7168_2048['B_w'].stride(1),
        _state_64_7168_2048['out'].stride(0), _state_64_7168_2048['out'].stride(1),
        _state_64_7168_2048['B_sc'].stride(0), _state_64_7168_2048['B_sc'].stride(1),
        BLOCK_SIZE_M=16, BLOCK_SIZE_N=128,
        BLOCK_SIZE_K=256,
        num_warps=4, num_stages=2,
        waves_per_eu=2, matrix_instr_nonkdim=16,
        cache_modifier=".cg",
    )
    return _state_64_7168_2048['out']


# ============================================================
# m=256, n=3072, k=1536
# ============================================================

"""M=256, N=3072, K=1536 ? native v_cvt_scalef32_pk_fp4_bf16 quant."""
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm

ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
MXFP4_QUANT_BLOCK_SIZE = 32

_state_256_3072_1536 = None


@triton.jit
def _mxfp4_quant_op_native_256_3072_1536(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    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)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    hw_scale = tl.exp2(scale_e8m0_unbiased)
    x_bf16 = x.to(tl.bfloat16)
    x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    x_lo, x_hi = tl.split(x_pairs)
    x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
    x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
    packed_bf16x2 = x_lo_i32 | x_hi_i32

    hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))

    hw_result = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v,v,v",
        [packed_bf16x2, hw_scale_bc],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )
    x_fp4 = (hw_result & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.heuristics({
    "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
    and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
})
@triton.jit
def _fused_mxfp4_quant_shuffle_kernel_256_3072_1536(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
    M, N,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr,
    SCALING_MODE: tl.constexpr, SCALE_N_PAD: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
        else:
            x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)

        out_tensor, bs_e8m0 = _mxfp4_quant_op_native_256_3072_1536(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, out_tensor, cache_modifier=".cg")
        else:
            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_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
        bs_offs_0 = bs_offs_m[:, None] // 32
        bs_offs_1 = bs_offs_m[:, None] % 32
        bs_offs_2 = bs_offs_1 % 16
        bs_offs_1 = bs_offs_1 // 16
        bs_offs_3 = bs_offs_n[None, :] // 8
        bs_offs_4 = bs_offs_n[None, :] % 8
        bs_offs_5 = bs_offs_4 % 4
        bs_offs_4 = bs_offs_4 // 4
        bs_offs = (bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 2 * 2
                   + bs_offs_5 * 2 * 2 * 16 + bs_offs_3 * 2 * 2 * 16 * 4
                   + bs_offs_0 * 2 * 16 * SCALE_N_PAD)
        bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
        bs_e8m0 = tl.where(bs_mask_valid, bs_e8m0, 127)
        SCALE_M_PAD = (M + 255) // 256 * 256
        bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[None, :]
        tl.store(bs_ptr + bs_offs, bs_e8m0.to(tl.uint8), mask=bs_mask, cache_modifier=".cg")


def _run_256_3072_1536(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    global _state_256_3072_1536
    if _state_256_3072_1536 is None:
        SCALE_N_valid = triton.cdiv(k, MXFP4_QUANT_BLOCK_SIZE)
        SCALE_M = triton.cdiv(m, 256) * 256
        SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
        BLOCK_SIZE_M = triton.cdiv(min(32, triton.next_power_of_2(m)), 32) * 32
        BLOCK_SIZE_N = 64
        grid = (triton.cdiv(m, BLOCK_SIZE_M), triton.cdiv(k, BLOCK_SIZE_N * 1))
        padded_M = (m + 31) // 32 * 32

        _state_256_3072_1536 = {
            'x_fp4': torch.empty((m, k // 2), dtype=torch.uint8, device=A.device),
            'blockscale': torch.empty((SCALE_M, SCALE_N), dtype=torch.uint8, device=A.device),
            'gemm_out': torch.empty((padded_M, n), dtype=torch.bfloat16, device=A.device),
            'SCALE_N': SCALE_N,
            'BLOCK_SIZE_M': BLOCK_SIZE_M,
            'BLOCK_SIZE_N': BLOCK_SIZE_N,
            'grid': grid,
        }

    s = _state_256_3072_1536

    _fused_mxfp4_quant_shuffle_kernel_256_3072_1536[s['grid']](
        A, s['x_fp4'], s['blockscale'],
        *A.stride(), *s['x_fp4'].stride(),
        M=m, N=k,
        BLOCK_SIZE_M=s['BLOCK_SIZE_M'], BLOCK_SIZE_N=s['BLOCK_SIZE_N'],
        NUM_ITER=1, NUM_STAGES=1,
        MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, SCALE_N_PAD=s['SCALE_N'],
        num_warps=2, waves_per_eu=0, num_stages=1,
    )

    gemm_a4w4_asm(
        s['x_fp4'].view(dtypes.fp4x2), B_shuffle,
        s['blockscale'].view(dtypes.fp8_e8m0), B_scale_sh,
        s['gemm_out'], ASM_KERNEL, None, 1.0, 0.0, True, log2_k_split=0,
    )

    return s['gemm_out'][:m]


import aiter as _aiter
from aiter import dtypes as _dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant as _dq
from aiter.utility.fp4_utils import e8m0_shuffle as _es
def _run_default(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
    A = A.contiguous()
    x_fp4, bs = _dq(A)
    bs = _es(bs)
    return _aiter.gemm_a4w4(x_fp4.view(_dtypes.fp4x2), B_shuffle, bs.view(_dtypes.fp8_e8m0), B_scale_sh, dtype=_dtypes.bf16, bpreshuffle=True)

from task import input_t, output_t
_DISPATCH = {
    (4, 2880, 512): _run_4_2880_512,
    (16, 2112, 7168): _run_16_2112_7168,
    (32, 4096, 512): _run_32_4096_512,
    (32, 2880, 512): _run_32_2880_512,
    (64, 7168, 2048): _run_64_7168_2048,
    (256, 3072, 1536): _run_256_3072_1536,
}
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n, _ = B.shape
    fn = _DISPATCH.get((m, n, k), _run_default)
    return fn(A, B, B_q, B_shuffle, B_scale_sh, m, n, k)
scrolls · 1020 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