Skip to content
KernelIndex
Search⌘K

submission 521876

rishi048401 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-521876?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
13.4µs
#430 of 1143
2026-03-09

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c7e3ca40db8bdcfea3a59749ef3e696d11c865fca1c2480bda4eb8d6705c778f
license declaredunknown
license concludedunknown
authorsrishi048401
imported2026-08-26

Kernel source

submission.py177 lines
try:
    from task import input_t, output_t
except ImportError:
    input_t = output_t = any

import torch
import triton
import triton.language as tl
from aiter.utility import dtypes
import aiter

@triton.jit
def my_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

    if SHUFFLE:
        # bitwise ALUs for massive optimization since AMD emulates integer division
        bs_offs_0 = bs_offs_m[:, None] >> 5         # // 32
        bs_offs_1 = bs_offs_m[:, None] & 31         # % 32
        bs_offs_2 = bs_offs_1 & 15                  # % 16
        bs_offs_1 = bs_offs_1 >> 4                  # // 16
        bs_offs_3 = bs_offs_n[None, :] >> 3         # // 8
        bs_offs_4 = bs_offs_n[None, :] & 7          # % 8
        bs_offs_5 = bs_offs_4 & 3                   # % 4
        bs_offs_4 = bs_offs_4 >> 2                  # // 4
        bs_offs = (
            bs_offs_1
            + (bs_offs_4 << 1)
            + (bs_offs_2 << 2)
            + (bs_offs_5 << 6)
            + (bs_offs_3 << 8)
            + bs_offs_0 * (32 * scaleN_pad)
        )
        bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
        bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
        bs_e8m0  = tl.where(bs_mask1, bs_e8m0, 127)
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
    else:
        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 custom_quant(x: torch.Tensor, config: dict):
    M, N = x.shape
    scaleN_valid = N // 32
    scaleM_pad   = (M + 31) // 32 * 32
    scaleN_pad   = (scaleN_valid +  7) //  8 *  8
    scaleN       = ((scaleN_valid + 7) // 8) * 8

    # Use tuned config or golden defaults
    b_size   = config.get("bs", 16 if M <= 16 else 64)
    n_warps  = config.get("w", 1 if M <= 16 else 4)
    n_stages = config.get("s", 1 if M <= 16 else 3)

    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
    bs    = torch.empty((((M + 255) // 256) * 256, scaleN), dtype=torch.uint8, device=x.device)
    grid  = ((M + b_size - 1) // b_size, scaleN)

    my_quant_kernel[grid](
        x, x_fp4, bs,
        N, 1,
        N // 2, 1,
        scaleN, 1,
        M, N,
        scaleN, scaleM_pad, scaleN_pad,
        BLOCK_SIZE=b_size,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0, SHUFFLE=True,
        num_warps=n_warps, num_stages=n_stages,
    )
    return x_fp4.view(dtypes.fp4x2), bs.view(dtypes.fp8_e8m0)


# ============================================================
# PHASE 5: SNIPER TUNING TABLE
# Config: {"bs": BLOCK_SIZE, "w": num_warps, "s": num_stages, "k": log2_k_split}
# ============================================================
CONFIG_TABLE = {
    (4, 2880):   {"bs": 16, "w": 1, "s": 1, "k": None},
    (8, 2112):   {"bs": 16, "w": 1, "s": 1, "k": None},
    (16, 2112):  {"bs": 16, "w": 1, "s": 1, "k": 1},
    (16, 3072):  {"bs": 16, "w": 1, "s": 1, "k": None},
    (32, 2880):  {"bs": 32, "w": 4, "s": 3, "k": None},
    (32, 4096):  {"bs": 32, "w": 4, "s": 3, "k": None},
    (64, 3072):  {"bs": 64, "w": 4, "s": 3, "k": None},
    (64, 7168):  {"bs": 64, "w": 4, "s": 3, "k": None},
    (256, 2880): {"bs": 64, "w": 4, "s": 3, "k": None},
    (256, 3072): {"bs": 64, "w": 4, "s": 3, "k": None},
}

_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNEL_64x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_OUT_CACHE = {}

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[0]

    # Shape-Specific Lookup
    config = CONFIG_TABLE.get((M, N), {})

    A_q, A_scale_sh = custom_quant(A, config)

    key = (M, N)
    if key not in _OUT_CACHE:
        _OUT_CACHE[key] = torch.empty((M, N), dtype=dtypes.bf16, device=A.device)
    out = _OUT_CACHE[key]

    kernel_name = _KERNEL_32x128 if M <= 64 else _KERNEL_64x128
    k_split = config.get("k", 1 if (M <= 16 and K >= 4096) else None)

    return aiter.gemm_a4w4_asm(
        A_q, B_shuffle, A_scale_sh, B_scale_sh, out,
        kernel_name, bpreshuffle=True, log2_k_split=k_split,
    )
scrolls · 177 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