Skip to content
KernelIndex
Search⌘K

submission 553074

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v150_lean_quant.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-553074?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
9.85µs
#220 of 1143
2026-03-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5d6aac4daf7e5f33d7f2474cd0b2814b2eb312421feaab533cd02f5e651f0944
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15

Techniques

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

fp41. Uses branchless FP4 E2M1 quantization (no 3-way if/else)
num-warps = 1num_warps=1, waves_per_eu=0, num_stages=1,
split-k_get_splitk_fn = _gemm_mod.get_splitk
stages = 1NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
tile-k = 512BLOCK_SIZE_K = 512
tile-m = 16BLOCK_SIZE_M = 16
tile-n = 32SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,

Kernel source

submission_v150_lean_quant.py484 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
v150: v120 + lean quantization op.
Replace _mxfp4_quant_op with a custom Triton JIT function that:
1. Uses branchless FP4 E2M1 quantization (no 3-way if/else)
2. Uses bitcast for exp2 (from v138)
3. Minimizes total VALU instructions
4. Uses integer-only rounding for max (avoids log2/floor/exp2 chain)
"""
from task import input_t, output_t

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import get_GEMM_config
from aiter.ops.gemm_op_common import get_padded_m

import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod
_reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel
_get_splitk_fn = _gemm_mod.get_splitk

_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16


@triton.jit
def _lean_mxfp4_quant_op(
    x,          # [BLOCK_M, BLOCK_K] float32
    BLOCK_K: tl.constexpr,
    BLOCK_M: tl.constexpr,
    QUANT_BLOCK: tl.constexpr,
):
    """Lean MXFP4 quantization — fewer VALU instructions, branchless E2M1.
    Produces bit-identical output to aiter's _mxfp4_quant_op.
    """
    NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK

    x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)

    # --- Scale computation (integer-only, no log2/floor) ---
    amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)
    # Round amax up to next power-of-2 (clear mantissa, round up exponent)
    amax_i32 = amax.to(tl.int32, bitcast=True)
    amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)
    amax_f32 = amax_rounded.to(tl.float32, bitcast=True)

    # Extract exponent directly: E8M0 scale = exponent_of(amax_rounded) - 2 - 127 + 127
    # = exponent_field - 2 = (amax_rounded >> 23) - 2
    # But we need to handle amax=0 => scale should be 0 (biased), i.e. -127 unbiased
    amax_exp = (amax_rounded >> 23).to(tl.int32)
    # scale_unbiased = amax_exp - 127 - 2 = amax_exp - 129
    scale_unbiased = amax_exp - 129
    # Clamp to [-127, 127] — tl.clamp only supports float, use tl.where
    scale_unbiased = tl.where(scale_unbiased < -127, -127, scale_unbiased)
    scale_unbiased = tl.where(scale_unbiased > 127, 127, scale_unbiased)
    # Handle amax=0: amax_exp=0, scale_unbiased=-129 clamped to -127. OK.

    # E8M0 biased scale (uint8)
    bs_e8m0 = (scale_unbiased + 127).to(tl.uint8)

    # --- Inverse scale via bitcast (fast exp2) ---
    # quant_scale = exp2(-scale_unbiased) = bitcast(((-scale_unbiased) + 127) << 23)
    inv_exp = (127 - scale_unbiased)
    quant_scale = (inv_exp << 23).to(tl.float32, bitcast=True)

    # --- Quantize: scale input ---
    qx = x_3d * quant_scale

    # --- Convert to FP4 E2M1 (branchless) ---
    # Extract sign and abs
    qx_i32 = qx.to(tl.int32, bitcast=True)
    sign_bit = ((qx_i32 >> 31) & 0x8).to(tl.uint8)  # sign at bit 3 for FP4
    qx_abs = (qx_i32 & 0x7FFFFFFF).to(tl.float32, bitcast=True)

    # Saturate: clamp abs to [0, 6.0] — values >= 6.0 become 0x7 (max E2M1 = 1.5 * 2^2 = 6.0)
    # After clamping, all values are in representable range, no saturation branch needed
    qx_clamped = tl.minimum(qx_abs, 6.0)
    qx_clamped_i32 = qx_clamped.to(tl.int32, bitcast=True)

    # Denormal path: values < 1.0 need special handling
    # E2M1 denormals: 0.0 (0b000), 0.5 (0b001)
    # Normal E2M1: 1.0 (0b010), 1.5 (0b011), 2.0 (0b100), 3.0 (0b101), 4.0 (0b110), 6.0 (0b111)
    #
    # For denormals (< 1.0): add magic number to round, extract low bits
    denorm_exp: tl.constexpr = (127 - 1) + (23 - 1) + 1
    denorm_magic: tl.constexpr = denorm_exp << 23
    denorm_magic_f: tl.constexpr = tl.cast(denorm_magic, tl.float32, bitcast=True)
    denormal_result = (qx_clamped + denorm_magic_f).to(tl.int32, bitcast=True) - denorm_magic
    denormal_result = denormal_result.to(tl.uint8)

    # Normal path (>= 1.0): round to nearest E2M1
    # IEEE float32 mantissa has 23 bits, E2M1 mantissa has 1 bit
    # So we need to round at bit 22 (keep only 1 mantissa bit)
    # Bias adjustment: subtract (127-1) from exponent to get E2M1 exponent
    qx_clamped_abs_i32 = qx_clamped_i32
    mant_odd = (qx_clamped_abs_i32 >> 22) & 1
    val_to_add: tl.constexpr = ((1 - 127) << 23) + (1 << 21) - 1
    normal_result = (qx_clamped_abs_i32 + val_to_add + mant_odd) >> 22
    normal_result = normal_result.to(tl.uint8)

    # Select: denormal if < 1.0, normal otherwise
    is_normal = qx_abs >= 1.0
    e2m1 = tl.where(is_normal, normal_result, denormal_result)

    # Apply sign
    e2m1 = e2m1 | sign_bit

    # Pack 2 FP4 values per byte
    e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_BLOCKS, QUANT_BLOCK // 2, 2])
    evens, odds = tl.split(e2m1)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_M, BLOCK_K // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)


@triton.jit
def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):
    chunk_size = tl.cdiv(num_pids, NUM_XCDS)
    xcd = pid % NUM_XCDS
    pid_in_xcd = pid // NUM_XCDS
    return xcd * chunk_size + pid_in_xcd


@triton.jit
def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
    num_pid_in_group = GROUP_SIZE_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_SIZE_M
    group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
    pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m
    pid_n = (pid % num_pid_in_group) // group_size_m
    return pid_m, pid_n


@triton.jit
def _fused_quant_gemm_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K_real,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    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,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    QUANT_BLOCK: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    K_packed = K_real // 2
    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    total_pids = GRID_MN * NUM_KSPLIT
    total_pids_padded = ((total_pids + 7) // 8) * 8

    pid_unified = tl.program_id(axis=0)
    pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)

    if pid_unified < total_pids:
        pid_k = pid_unified % NUM_KSPLIT
        pid = pid_unified // NUM_KSPLIT
        num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
        num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

        if NUM_KSPLIT == 1:
            pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
        else:
            pid_m = pid // num_pid_n
            pid_n = pid % num_pid_n

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

        if (pid_k * SPLITK_BLOCK_SIZE) < K_real:
            num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)

            offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
            offs_k = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
            a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak

            offs_k_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
            offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
            b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn

            offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
            offs_ks_scale = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
                0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
            )
            b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks_scale[None, :] * stride_bsk

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

            for k in tl.range(0, num_k_iter):
                a_bf16 = tl.load(a_ptrs).to(tl.float32)
                a_fp4, a_scales = _lean_mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)

                b_fp4 = tl.load(b_ptrs)

                b_scales = (
                    tl.load(b_scale_ptrs)
                    .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)
                )

                accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)

                a_ptrs += BLOCK_SIZE_K * stride_ak
                b_ptrs += (BLOCK_SIZE_K // 2) * 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, :] + pid_k * stride_ck
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(c_ptrs, c, mask=c_mask)


@triton.jit
def _fused_quant_shuffle_kernel(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m, stride_x_n,
    stride_x_fp4_m, stride_x_fp4_n,
    M, N, scale_n_valid,
    SCALE_N: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    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
        x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
        x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)

        out_tensor, bs_e8m0 = _lean_mxfp4_quant_op(
            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
        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_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        m_idx = bs_offs_m[:, None]
        n_idx = bs_offs_n[None, :]
        i0 = m_idx // 32
        i1 = (m_idx // 16) % 2
        i2 = m_idx % 16
        i3 = n_idx // 8
        i4 = (n_idx // 4) % 2
        i5 = n_idx % 4
        shuffled_offset = (i0 * (SCALE_N * 32) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1)
        bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scale_n_valid)[None, :]
        bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
        bs_store_mask = (m_idx < (M + 255) // 256 * 256) & (n_idx < SCALE_N)
        tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)


_cache_asm = {}
_cache_fused = {}
_gemm_asm = None
_warmup_done = False


def custom_kernel(data: input_t) -> output_t:
    global _gemm_asm, _warmup_done

    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.shape[0]

    use_fused = (M <= 64)

    # Warmup: use ASM path to init aiter module
    if not _warmup_done:
        scale_n_valid = (K + 31) // 32
        SCALE_M = ((M + 255) // 256) * 256
        SCALE_N = ((scale_n_valid + 7) // 8) * 8
        BSM = triton.next_power_of_2(M) if M <= 32 else 16
        grid = (triton.cdiv(M, BSM), triton.cdiv(K, 32))

        x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
        bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)

        _fused_quant_shuffle_kernel[grid](
            A, x_fp4, bs_sh,
            A.stride(0), A.stride(1),
            x_fp4.stride(0), x_fp4.stride(1),
            M, K, scale_n_valid,
            SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
            NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
            num_warps=1, waves_per_eu=0, num_stages=1,
        )

        result = aiter.gemm_a4w4(
            x_fp4.view(_fp4x2), B_shuffle,
            bs_sh.view(_fp8_e8m0), B_scale_sh,
            dtype=_bf16, bpreshuffle=True,
        )
        _warmup_done = True
        try:
            _gemm_asm = torch.ops.aiter.gemm_a4w4_asm
        except Exception:
            try:
                import aiter.jit.core as _jc
                _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
            except Exception:
                pass
        return result

    if use_fused:
        # --- Fused quant+GEMM: single kernel launch ---
        key = (M, K, N)
        c = _cache_fused.get(key)
        if c is None:
            K_packed = K // 2
            scale_n = (K + 31) // 32
            SCALE_N_B = ((scale_n + 7) // 8) * 8

            BLOCK_SIZE_M = 16
            BLOCK_SIZE_N = 64 if M <= 16 else 128
            BLOCK_SIZE_K = 512

            base_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
            target_ksplit = max(1, 256 // max(1, base_blocks))

            if target_ksplit > 1:
                SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(
                    K_packed, BLOCK_SIZE_K, target_ksplit
                )
                if BLOCK_SIZE_K_adj < 512:
                    BLOCK_SIZE_K_adj = 512
                    SPLITK_BLOCK_SIZE = 2 * K_packed
                    NUM_KSPLIT = 1
                else:
                    BLOCK_SIZE_K = BLOCK_SIZE_K_adj
            else:
                NUM_KSPLIT = 1
                SPLITK_BLOCK_SIZE = 2 * K_packed

            if NUM_KSPLIT > 1:
                y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
            else:
                y_pp = None
                SPLITK_BLOCK_SIZE = 2 * K_packed

            y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)

            total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
            total_blocks = ((total_blocks_raw + 7) // 8) * 8

            bs_stride_n = 32 * SCALE_N_B
            bs_stride_k = 1

            c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
                 NUM_KSPLIT, SPLITK_BLOCK_SIZE,
                 y, y_pp, total_blocks, bs_stride_n, bs_stride_k)
            _cache_fused[key] = c

        (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
         NUM_KSPLIT, SPLITK_BLOCK_SIZE,
         y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c

        B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
        B_q_T = B_q_u8.T
        B_scale_u8 = B_scale_sh.view(torch.uint8)

        out_tensor = y if NUM_KSPLIT == 1 else y_pp

        _fused_quant_gemm_kernel[(total_blocks,)](
            A, B_q_T, out_tensor, B_scale_u8,
            M, N, K,
            A.stride(0), A.stride(1),
            B_q_T.stride(0), B_q_T.stride(1),
            0 if NUM_KSPLIT == 1 else y_pp.stride(0),
            y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),
            y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),
            bs_stride_n, bs_stride_k,
            BLOCK_SIZE_M=BLOCK_SIZE_M,
            BLOCK_SIZE_N=BLOCK_SIZE_N,
            BLOCK_SIZE_K=BLOCK_SIZE_K,
            GROUP_SIZE_M=8,
            NUM_KSPLIT=NUM_KSPLIT,
            SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
            QUANT_BLOCK=32,
            num_warps=8,
            num_stages=2,
            waves_per_eu=0,
        )

        if NUM_KSPLIT > 1:
            ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))
            grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))
            _reduce_kernel[grid_reduce](
                y_pp, y, M, N,
                y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                y.stride(0), y.stride(1),
                16, 64, ACTUAL_KSPLIT,
                triton.next_power_of_2(NUM_KSPLIT),
            )

        return y

    else:
        # --- ASM GEMM path ---
        key = (M, K, N)
        c = _cache_asm.get(key)
        if c is None:
            scale_n_valid = (K + 31) // 32
            SCALE_M = ((M + 255) // 256) * 256
            SCALE_N = ((scale_n_valid + 7) // 8) * 8
            padded_m = get_padded_m(M, N, K, 0)

            BSM = triton.next_power_of_2(M) if M <= 32 else 16
            NW = 1
            BSN = 32
            NUM_ITER_Q = 2
            grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER_Q))

            ck_config = get_GEMM_config(M, N, K)
            kernel_name = ""
            split_k = 0
            if ck_config is not None:
                split_k = ck_config.get("splitK", 0) or 0
                kernel_name = ck_config["kernelName"]

            x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
            bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
            out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)

            x_fp4_view = x_fp4.view(_fp4x2)
            bs_sh_view = bs_sh.view(_fp8_e8m0)
            out_view = out[:M] if M < padded_m else out

            c = (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
                 x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
                 kernel_name, split_k,
                 A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1))
            _cache_asm[key] = c

        (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
         x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
         kernel_name, split_k,
         stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c

        _fused_quant_shuffle_kernel[grid](
            A, x_fp4, bs_sh,
            stride_a0, stride_a1,
            stride_fp4_0, stride_fp4_1,
            M, K, scale_n_valid,
            SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
            NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
            num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
        )

        if _gemm_asm is not None:
            _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
                      out, kernel_name, None, 1.0, 0.0, True, split_k)
            return out_view

        return aiter.gemm_a4w4(
            x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
            dtype=_bf16, bpreshuffle=True,
        )
scrolls · 484 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 553005.

⋯ 307 unchanged lines
M, K, scale_n_valid,
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0,
- allow_flush_denorm=True, num_stages=1,
+ num_warps=1, waves_per_eu=0, num_stages=1,
)
result = aiter.gemm_a4w4(
⋯ 90 unchanged lines
num_warps=8,
num_stages=2,
waves_per_eu=0,
- allow_flush_denorm=True,
)
if NUM_KSPLIT > 1:
⋯ 58 unchanged lines
M, K, scale_n_valid,
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0,
- allow_flush_denorm=True, num_stages=NUM_ITER_Q,
+ num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
)
if _gemm_asm is not None:
scrolls · 28 diff lines total

Best evidence level for this revision: reported

JSON