Skip to content
KernelIndex
Search⌘K

submission 754001

Navdeep Singh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754001?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
63.0µs
#1132 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:080933fe27ec1af851bd56aa8db08e595702acf1cff5a246174e2aba03e8e4d8
license declaredunknown
license concludedunknown
authorsNavdeep Singh
imported2026-08-26

Techniques

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

fp8return e5m2_bytes.to(tl.int8).to(tl.float8e5, bitcast=True)
mmainner_acc = tl.dot(a_fp8, tl.trans(b_fp8), out_dtype=tl.float32)
num-warps = 16block_m, block_n, block_k, num_warps = 16, 64, 32, 4
split-kGROUP_M: tl.constexpr, SPLIT_K: tl.constexpr, EVEN_K: tl.constexpr,

Kernel source

submission.py200 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant

def _inverse_e8m0_shuffle(scale_shuffled: torch.Tensor, logical_m: int) -> torch.Tensor:
    scale_u8 = scale_shuffled.view(torch.uint8)
    sm, sn = scale_u8.shape
    scale = (
        scale_u8.view(sm // 32, sn // 8, 4, 16, 2, 2)
        .permute(0, 5, 3, 1, 4, 2)
        .contiguous()
        .view(sm, sn)
    )
    return scale[:logical_m].contiguous()


@triton.jit
def _pid_grid(pid: int, num_pid_m: int, num_pid_n: int, 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
    actual_group_size_m = tl.minimum(num_pid_m - first_pid_m, group_size_m)
    pid_m = first_pid_m + (pid % actual_group_size_m)
    pid_n = (pid % num_pid_in_group) // actual_group_size_m
    return pid_m, pid_n


@triton.jit
def decode_e2m1_to_fp8(packed_bytes, is_odd):
    # Natively transforms MXFP4 e2m1 into standard fp8 native types using pure logic
    nibbles = tl.where(is_odd, packed_bytes >> 4, packed_bytes & 0xF).to(tl.int32)
    abs_val = nibbles & 0x7
    base = 56 + (abs_val << 1)
    res = tl.where(abs_val == 1, 56, base)
    res = tl.where(abs_val == 0, 0, res)
    e5m2_bytes = res | ((nibbles & 0x8) << 4)
    return e5m2_bytes.to(tl.int8).to(tl.float8e5, bitcast=True)


@triton.heuristics({
    "EVEN_K": lambda args: args["K"] % args["BLOCK_K"] == 0,
})
@triton.jit
def _mxfp4_gemm_kernel(
    a_ptr, b_ptr, c_ptr,
    a_scales_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    stride_asm, stride_ask,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr, SPLIT_K: tl.constexpr, EVEN_K: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    pid_k = tl.program_id(axis=1)

    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_M)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    
    offs_k = tl.arange(0, BLOCK_K)
    byte_offs_k = offs_k // 2
    is_odd = (offs_k % 2) == 1

    a_ptrs = a_ptr + offs_m[:, None] * stride_am + byte_offs_k[None, :] * stride_ak
    b_ptrs = b_ptr + offs_n[:, None] * stride_bn + byte_offs_k[None, :] * stride_bk

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    k_tiles_per_split = tl.cdiv(K, BLOCK_K)
    if SPLIT_K > 1:
        tiles_per_split = tl.cdiv(k_tiles_per_split, SPLIT_K)
        start_k = pid_k * tiles_per_split
        end_k = tl.minimum((pid_k + 1) * tiles_per_split, k_tiles_per_split)
    else:
        start_k = 0
        end_k = k_tiles_per_split

    a_ptrs += start_k * (BLOCK_K // 2) * stride_ak
    b_ptrs += start_k * (BLOCK_K // 2) * stride_bk

    for k_i in range(start_k, end_k):
        # Scale loading (BLOCK_K is exactly 32)
        # Therefore each iteration handles exactly ONE scale
        a_scale_ptrs = a_scales_ptr + offs_m[:, None] * stride_asm + k_i * stride_ask
        b_scale_ptrs = b_scales_ptr + offs_n[:, None] * stride_bsn + k_i * stride_bsk

        if EVEN_K:
            a_bytes = tl.load(a_ptrs, mask=(offs_m[:, None] < M), other=0)
            b_bytes = tl.load(b_ptrs, mask=(offs_n[:, None] < N), other=0)
            
            a_scale_val = tl.load(a_scale_ptrs, mask=(offs_m[:, None] < M), other=0)
            b_scale_val = tl.load(b_scale_ptrs, mask=(offs_n[:, None] < N), other=0)
        else:
            k_mask_byte = (k_i * (BLOCK_K // 2) + byte_offs_k) < (K // 2)
            
            a_bytes = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & k_mask_byte[None, :], other=0)
            b_bytes = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & k_mask_byte[None, :], other=0)
            
            a_scale_val = tl.load(a_scale_ptrs, mask=(offs_m[:, None] < M), other=0)
            b_scale_val = tl.load(b_scale_ptrs, mask=(offs_n[:, None] < N), other=0)

        # Decode via bitcast algebraic logic perfectly fitting into registers!
        a_fp8 = decode_e2m1_to_fp8(a_bytes, is_odd[None, :])
        b_fp8 = decode_e2m1_to_fp8(b_bytes, is_odd[None, :])

        inner_acc = tl.dot(a_fp8, tl.trans(b_fp8), out_dtype=tl.float32)

        # Scale outer product algebraically
        a_sc_f32 = tl.exp2(a_scale_val.to(tl.float32) - 127.0) # [BLOCK_M, 1]
        b_sc_f32 = tl.exp2(b_scale_val.to(tl.float32) - 127.0) # [BLOCK_N, 1]
        
        acc += inner_acc * a_sc_f32 * tl.trans(b_sc_f32)

        a_ptrs += (BLOCK_K // 2) * stride_ak
        b_ptrs += (BLOCK_K // 2) * stride_bk

    acc_out = acc.to(c_ptr.type.element_ty)
    if SPLIT_K > 1:
        c_ptrs = c_ptr + pid_k * (M * N) + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    else:
        c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, acc_out, mask=c_mask)


def get_tune_config(M, N, K):
    if M <= 16:
        block_m, block_n, block_k, num_warps = 16, 64, 32, 4
        split_k = 16 if K >= 4096 else 8
    elif M <= 32:
        block_m, block_n, block_k, num_warps = 32, 64, 32, 4
        split_k = 8 if K >= 2048 else 4
    elif M <= 64:
        block_m, block_n, block_k, num_warps = 64, 64, 32, 4
        split_k = 4 if K >= 2048 else 2
    else:
        block_m, block_n, block_k, num_warps = 128, 128, 32, 8
        split_k = 1

    k_tiles = triton.cdiv(K, block_k)
    if k_tiles < split_k:
        split_k = max(1, k_tiles)

    return block_m, block_n, block_k, split_k, num_warps, 2


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data

    A = A.contiguous()

    M, K = A.shape
    N = B_q.shape[0]

    A_q_u8, A_scale_raw = dynamic_mxfp4_quant(A)
    # Native Linear Math arrays
    A_q = A_q_u8[:M, :K // 2].contiguous().view(torch.uint8)
    A_scale = A_scale_raw[:M, :K // 32].contiguous().view(torch.uint8)

    B_q_u8 = B_q.contiguous().view(torch.uint8)
    B_scale = _inverse_e8m0_shuffle(B_scale_sh, N).view(torch.uint8).contiguous()

    block_m, block_n, block_k, split_k, num_warps, num_stages = get_tune_config(M, N, K)

    if split_k > 1:
        out = torch.empty((split_k, M, N), device=A.device, dtype=torch.float32)
    else:
        out = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)

    grid = (triton.cdiv(M, block_m) * triton.cdiv(N, block_n), split_k)

    _mxfp4_gemm_kernel[grid](
        A_q, B_q_u8, out,
        A_scale, B_scale,
        M, N, K,
        A_q.stride(0), A_q.stride(1),
        B_q_u8.stride(0), B_q_u8.stride(1),
        out.stride(-2), out.stride(-1),
        A_scale.stride(0), A_scale.stride(1),
        B_scale.stride(0), B_scale.stride(1),
        BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k,
        GROUP_M=4, SPLIT_K=split_k,
        num_warps=num_warps, num_stages=num_stages
    )

    if split_k > 1:
        return out.sum(dim=0).to(torch.bfloat16)
    return out
scrolls · 200 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