Skip to content
KernelIndex
Search⌘K

submission 745344

ykaitao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-745344?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
#219 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:205bc0bc1762058f9ce739069421e31aad9e1b9e7a6f6c543bdb90cf743473d4
license declaredunknown
license concludedunknown
authorsykaitao
imported2026-08-15

Techniques

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

fp4"""Optimized MXFP4 quant: replaces log2/exp2 with bitwise float32 exponent ops."""
num-warps = 1NUM_WARPS = 1
split-kget_splitk,
stages = 1NUM_STAGES = 1
tile-k = 512_BK = 512
tile-m = 64BLOCK_M = 64
tile-n = 32BLOCK_N = 32

Kernel source

submission.py790 lines
"""v137: s6 log2_k_split=2 fix (was 0, causing 0.63 waves; now 2.53 waves).

Key changes vs v122:
- s6 (m=256,k=1536): log2_k_split was 0 (192 CTAs, 0.63 waves) → now 2 (768 CTAs, 2.53 waves)
- s5 (m=64,k=2048): quant+ASM with log2_k_split=2 → 448 blocks (1.47 waves) [unchanged]
- s2 (m=16,k=7168): bm=8,nsplit=8 → 528 blocks (1.74 waves) [unchanged]
- t2 (m=16,k=1536): preshuffle (bm=16,bn=64,nsplit=8 → 384 blocks) [unchanged]
- t3 (m=64,k=1536): preshuffle (bm=16,bn=128,nsplit=4 → 384 blocks) [unchanged]
"""

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

from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
    get_splitk,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op


@triton.jit
def _mxfp4_quant_op_fast(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE):
    """Optimized MXFP4 quant: replaces log2/exp2 with bitwise float32 exponent ops."""
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1

    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)

    # Calculate scale using bitwise ops (no log2/exp2 transcendental functions)
    amax_f32 = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax_i32 = amax_f32.to(tl.int32, bitcast=True)
    amax_u32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    # Extract biased exponent from float32 bits (bits 30..23)
    biased_exp = tl.cast((amax_u32 >> 23) & 0xFF, tl.int32)
    # bs_e8m0 = clamp(biased_exp - 2, 0, 254)  equivalent to original formula
    bs_e8m0_int = tl.maximum(tl.minimum(biased_exp - 2, 254), 0)
    bs_e8m0 = bs_e8m0_int.to(tl.uint8)
    # quant_scale = 2^(127 - bs_e8m0): float32 with biased_exp=(254-bs_e8m0), mantissa=0
    quant_scale_u32 = tl.cast(254 - bs_e8m0_int, tl.uint32) << 23
    quant_scale = quant_scale_u32.to(tl.float32, bitcast=True)

    qx = x * quant_scale
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    qx = qx ^ s
    qx_fp32 = qx.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = not (saturate_mask | denormal_mask)

    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)

    normal_x = qx
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)

    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp
    e2m1_value = tl.reshape(
        e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)

_BK = 512
_BUF: dict = {}
_PARAMS: dict = {}
_PARTIAL: dict = {}
_RUNNER: dict = {}
_RUNNER_ARGS: dict = {}
_REDUCE_RUNNER: dict = {}
_REDUCE_RUNNER_ARGS: dict = {}
_STATES: dict = {}
_BK_CACHE: dict = {}
_BS_CACHE: dict = {}
_DISPATCH_CACHE: dict = {}


def _use_preshuffle(M: int, N: int, K: int) -> bool:
    # s1/s3/s4 (K≤512), s2 (K=7168,M≤16): preshuffle avoids A-quant entirely
    # t2 (M=16,K=1536) and t3 (M=64,K=1536): route to preshuffle (fast path for small M)
    # s5 (M=64,K=2048) and s6 (M=256,K=1536): use quant+ASM for better wave utilization
    return (
        K <= 512
        or (K == 7168 and M <= 16)
        or (K == 1536 and M <= 64)
    )


def _unwrap_jit_fn(jit_fn):
    while not hasattr(jit_fn, "device_caches"):
        jit_fn = jit_fn.fn
    return jit_fn


def _get_kernel_cache(jit_fn):
    try:
        inner = _unwrap_jit_fn(jit_fn)
        device = torch.cuda.current_device()
        kernel_cache = inner.device_caches[device][0]
        return kernel_cache, device
    except Exception:
        return None, None


def _extract_runner(kernel_cache, keys_before, grid):
    try:
        new_keys = set(kernel_cache.keys()) - keys_before
        if not new_keys:
            return None
        compiled_kernel = kernel_cache[next(iter(new_keys))]
        compiled_kernel._init_handles()
        if len(grid) == 1:
            grid = (grid[0], 1, 1)
        elif len(grid) == 2:
            grid = (grid[0], grid[1], 1)
        return compiled_kernel[grid]
    except Exception:
        return None


def _check_even_k(K_int: int, bk: int, spk_sz: int) -> bool:
    return K_int % (bk // 2) == 0 and spk_sz % bk == 0 and K_int % (spk_sz // 2) == 0


def _shape_config(M: int, N: int = 0, K_bf16: int = 0) -> dict:
    # Shape2: bm=8,bn=64,nsplit=8 → 528 blocks (1.74 waves): confirmed 13.7µs vs 15.8µs (bm=16)
    if N == 2112 and K_bf16 == 7168 and M <= 16:
        return {
            "bm": 8,
            "bn": 64,
            "num_warps": 4,
            "waves_per_eu": 4,
            "cache_modifier": ".cg",
            "num_stages": 2,
        }
    # t2: m=16,n=3072,k=1536 → bm=16,bn=64,nsplit=8 → 384 blocks (1.26 waves)
    if M <= 16 and K_bf16 == 1536:
        return {
            "bm": 16,
            "bn": 64,
            "num_warps": 4,
            "waves_per_eu": 4,
            "cache_modifier": ".cg",
        }
    # t3: m=64,n=3072,k=1536 → bm=16,bn=128,nsplit=4 → 384 blocks (1.26 waves)
    if M <= 64 and K_bf16 == 1536:
        return {
            "bm": 16,
            "bn": 128,
            "num_warps": 4,
            "waves_per_eu": 4,
            "cache_modifier": ".cg",
        }
    # Shape5: m=64, k=2048 → quant+ASM path, shape config for preshuffle tile (unused for s5)
    if M <= 64 and K_bf16 == 2048:
        return {
            "bm": 16,
            "bn": 128,
            "num_warps": 4,
            "waves_per_eu": 2,
            "cache_modifier": ".cg",
        }
    # Shape6: m=256, k=1536 → quant+ASM path, shape config for preshuffle tile (unused for s6)
    if M == 256 and K_bf16 == 1536:
        return {
            "bm": 16,
            "bn": 128,
            "num_warps": 4,
            "waves_per_eu": 2,
            "cache_modifier": ".cg",
        }
    # Shape1: m=4, k=512 → bm=4, bn=32 (90 blocks, 30% CU) + .cg + waves=4
    if M <= 8 and K_bf16 <= 512:
        return {
            "bm": 4,
            "bn": 32,
            "num_warps": 2,
            "waves_per_eu": 4,
            "cache_modifier": ".cg",
        }
    # Shapes 3,4: m=32, k=512 → bm=16, bn=32 (s3: 256 blocks, s4: 180 blocks) + .cg + waves=4
    if M <= 32 and K_bf16 <= 512:
        return {
            "bm": 16,
            "bn": 32,
            "num_warps": 4,
            "waves_per_eu": 4,
            "cache_modifier": ".cg",
        }
    return {
        "bm": 16,
        "bn": 64,
        "num_warps": 4,
        "waves_per_eu": 2,
        "cache_modifier": None,
    }


def _splitk_target(M: int, N: int, K_bf16: int) -> int:
    # Shape2: use 8 splits → 528 blocks (1.74 waves) with bm=8
    if N == 2112 and K_bf16 == 7168 and M <= 16:
        return 8
    # t2: m=16,k=1536 → nsplit=8 → 384 blocks (1.26 waves)
    if K_bf16 == 1536 and M <= 16:
        return 8
    # t3: m=64,k=1536 → nsplit=4 → 384 blocks (1.26 waves) vs 96 blocks (0.32w) without
    if K_bf16 == 1536 and M <= 64:
        return 4
    return 1


def _build_params(M: int, N: int, K_bf16: int) -> dict:
    cfg = _shape_config(M, N, K_bf16)
    K_int = K_bf16 // 2
    num_ksplit = _splitk_target(M, N, K_bf16)
    if num_ksplit > 1:
        spk_sz, bk, num_ksplit = get_splitk(K_int, _BK, num_ksplit)
    else:
        spk_sz = 2 * K_int
        bk = triton.next_power_of_2(2 * K_int) if _BK >= 2 * K_int else _BK
    g0 = triton.cdiv(M, cfg["bm"]) * triton.cdiv(N, cfg["bn"])
    grid = (g0 * num_ksplit,)
    actual_ksplit = triton.cdiv(K_int, spk_sz // 2)
    max_ksplit = triton.next_power_of_2(num_ksplit)
    reduce_grid = (triton.cdiv(M, 16), triton.cdiv(N, 16)) if num_ksplit > 1 else None
    return {
        "bm": cfg["bm"],
        "bn": cfg["bn"],
        "num_warps": cfg["num_warps"],
        "waves_per_eu": cfg["waves_per_eu"],
        "cache_modifier": cfg["cache_modifier"],
        "bk": bk,
        "spk_sz": spk_sz,
        "grid": grid,
        "g0": g0,
        "K_int": K_int,
        "num_ksplit": num_ksplit,
        "actual_ksplit": actual_ksplit,
        "max_ksplit": max_ksplit,
        "reduce_grid": reduce_grid,
    }


def _build_main_args(p, A, B_k, C_out, B_s):
    K_int = p["K_int"]
    EVEN_K = _check_even_k(K_int, p["bk"], p["spk_sz"])
    GRID_MN = p["g0"]
    if p["num_ksplit"] == 1:
        stride_ck = 0
        stride_cm = C_out.stride(0)
        stride_cn = C_out.stride(1)
    else:
        stride_ck = C_out.stride(0)
        stride_cm = C_out.stride(1)
        stride_cn = C_out.stride(2)
    return [
        A,
        B_k,
        C_out,
        B_s,
        p["M_val"],
        p["N_val"],
        K_int,
        A.stride(0),
        A.stride(1),
        B_k.stride(0),
        B_k.stride(1),
        stride_ck,
        stride_cm,
        stride_cn,
        B_s.stride(0),
        B_s.stride(1),
        p["bm"],
        p["bn"],
        p["bk"],
        1,
        p["num_ksplit"],
        p["spk_sz"],
        EVEN_K,
        p["num_warps"],
        p.get("num_stages", 1),
        p["waves_per_eu"],
        16,
        GRID_MN,
        True,
        p["cache_modifier"],
    ]


def _build_reduce_args(p, y_pp, C):
    return [
        y_pp,
        C,
        p["M_val"],
        p["N_val"],
        y_pp.stride(0),
        y_pp.stride(1),
        y_pp.stride(2),
        C.stride(0),
        C.stride(1),
        16,
        16,
        p["actual_ksplit"],
        p["max_ksplit"],
    ]


@triton.jit
def _mxfp4_quant_shuffled_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_shuffled_ptr,
    stride_x_m,
    stride_x_n,
    stride_fp4_m,
    stride_fp4_n,
    M,
    N,
    SN,
    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,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER

    stride_x_m = tl.cast(stride_x_m, tl.int64)
    stride_x_n = tl.cast(stride_x_n, tl.int64)
    stride_fp4_m = tl.cast(stride_fp4_m, tl.int64)
    stride_fp4_n = tl.cast(stride_fp4_n, tl.int64)
    SN64 = tl.cast(SN, tl.int64)

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, start_n + NUM_ITER, 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, other=0.0, cache_modifier=".cg"
            ).to(tl.float32)

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

        fp4_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        fp4_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        fp4_offs = (
            fp4_offs_m[:, None] * stride_fp4_m + fp4_offs_n[None, :] * stride_fp4_n
        )
        if EVEN_M_N:
            tl.store(x_fp4_ptr + fp4_offs, out_tensor)
        else:
            fp4_mask = (fp4_offs_m < M)[:, None] & (fp4_offs_n < (N // 2))[None, :]
            tl.store(x_fp4_ptr + fp4_offs, out_tensor, mask=fp4_mask)

        x_idx = tl.cast(pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M), tl.int64)
        y_idx = tl.cast(
            pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS), tl.int64
        )
        xr = x_idx[:, None]
        yr = y_idx[None, :]
        sh_idx = (
            (xr // 32 * SN64) * 32
            + (yr // 8) * 256
            + (yr % 4) * 64
            + (xr % 16) * 4
            + (yr % 8) // 4 * 2
            + (xr % 32) // 16
        )
        if EVEN_M_N:
            tl.store(bs_shuffled_ptr + sh_idx, bs_e8m0)
        else:
            bs_mask = (x_idx[:, None] < M) & (
                y_idx[None, :]
                < (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            )
            tl.store(bs_shuffled_ptr + sh_idx, bs_e8m0, mask=bs_mask)


def _quant_config(M: int, K: int) -> dict:
    if M <= 32:
        BLOCK_M = triton.next_power_of_2(M)
        BLOCK_N = 32
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 1
    else:
        BLOCK_M = 64
        BLOCK_N = 64
        NUM_ITER = 4
        NUM_STAGES = 2
        NUM_WARPS = 4
        if K <= 16384:
            # s5 (M=64, K=2048): BLOCK_M=8 → 256 blocks (84%, 1-wave) vs 128 blocks (42%)
            # s6 (M=256, K=1536): BLOCK_M=8 → 768 blocks (2.5 waves), keep 16
            BLOCK_M = 8 if M <= 64 else 16
            BLOCK_N = 64
            NUM_ITER = 1

    if K <= 1024:
        BLOCK_N = min(256, triton.next_power_of_2(K))
        BLOCK_N = max(32, BLOCK_N)
        BLOCK_M = min(8, triton.next_power_of_2(M))
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4

    EVEN_M_N = (M % BLOCK_M == 0) and (K % (BLOCK_N * NUM_ITER) == 0)
    grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_N * NUM_ITER))
    return {
        "BLOCK_M": BLOCK_M,
        "BLOCK_N": BLOCK_N,
        "NUM_ITER": NUM_ITER,
        "NUM_STAGES": NUM_STAGES,
        "NUM_WARPS": NUM_WARPS,
        "EVEN_M_N": EVEN_M_N,
        "grid": grid,
    }


def _select_gemm_kernel(M: int, N: int, K: int) -> str:
    """Select optimal ASM GEMM kernel name for register efficiency.

    The 32x128 tile gives better occupancy than 192x128 when M is not a multiple
    of 192. All benchmark shapes (M=64, M=256) benefit from 32x128 tile:
    - M=64: 32x128 → 2x56=112 blocks vs 1x56=56 blocks (2x better occupancy)
    - M=256: 32x128 → 8x24=192 blocks vs 2x24=48 blocks (4x better occupancy)
    Use 32x128 for all shapes in quant+ASM path since register utilization is better.
    """
    return "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"


def _init_state(M: int, N: int, K: int, device) -> dict:
    SM = (M + 255) // 256 * 256
    SN = (K // 32 + 7) // 8 * 8
    padded_M = (M + 31) // 32 * 32

    A_q_buf = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
    A_scale_sh_buf = torch.zeros((SM, SN), dtype=torch.uint8, device=device)
    out_buf = torch.empty((padded_M, N), dtype=torch.bfloat16, device=device)

    A_q_fp4 = A_q_buf.view(dtypes.fp4x2)
    A_scale_fp8 = A_scale_sh_buf.view(dtypes.fp8_e8m0)
    cfg = _quant_config(M, K)
    gemm_kernel_name = _select_gemm_kernel(M, N, K)

    return {
        "A_q_buf": A_q_buf,
        "A_scale_sh_buf": A_scale_sh_buf,
        "out_buf": out_buf,
        "A_q_fp4": A_q_fp4,
        "A_scale_fp8": A_scale_fp8,
        "SM": SM,
        "SN": SN,
        "padded_M": padded_M,
        "cfg": cfg,
        "gemm_kernel_name": gemm_kernel_name,
        "quant_runner": None,
        "quant_ra": None,
        # log2_k_split: s5 (K=2048,M<=64) → lk=2 (768 CTAs, 2.53w), s6 (K=1536,M=256) → lk=2 (768 CTAs, 2.53w)
        "log2_k_split": 2 if (K == 2048 and M <= 64) else (2 if (K == 1536 and M == 256) else 0),
    }


def _build_quant_args(cfg, SN, A, A_q_buf, A_scale_sh_buf, M, K) -> list:
    return [
        A,
        A_q_buf,
        A_scale_sh_buf,
        A.stride(0),
        A.stride(1),
        A_q_buf.stride(0),
        A_q_buf.stride(1),
        M,
        K,
        SN,
        cfg["BLOCK_M"],
        cfg["BLOCK_N"],
        cfg["NUM_ITER"],
        cfg["NUM_STAGES"],
        32,
        cfg["EVEN_M_N"],
    ]


def _custom_kernel_preshuffle(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
    M: int,
    N: int,
    K_bf16: int,
) -> torch.Tensor:
    shape_key = (M, N, K_bf16)
    K_packed = B_shuffle.shape[1]
    SN_B = B_scale_sh.shape[1]

    C = _BUF.get(shape_key)
    if C is None:
        C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        _BUF[shape_key] = C

    p = _PARAMS.get(shape_key)
    if p is None:
        p = _build_params(M, N, K_bf16)
        p["M_val"] = M
        p["N_val"] = N
        _PARAMS[shape_key] = p

    y_pp = _PARTIAL.get(shape_key)
    if p["num_ksplit"] > 1 and y_pp is None:
        y_pp = torch.empty(
            (p["num_ksplit"], M, N), dtype=torch.float32, device=A.device
        )
        _PARTIAL[shape_key] = y_pp

    # Always recompute B_k/B_s views from current B tensors (B changes with each seed)
    B_k = torch.as_strided(
        B_shuffle.view(torch.uint8),
        (N // 16, K_packed * 16),
        (K_packed * 16, 1),
    )
    B_s = torch.as_strided(
        B_scale_sh.view(torch.uint8),
        (N // 32, K_bf16),
        (SN_B * 32, 1),
    )

    runner = _RUNNER.get(shape_key)
    reduce_runner = _REDUCE_RUNNER.get(shape_key)
    C_out = C if p["num_ksplit"] == 1 else y_pp

    if runner is None:
        main_kc, _ = _get_kernel_cache(_gemm_a16wfp4_preshuffle_kernel)
        main_keys_before = set(main_kc.keys()) if main_kc is not None else None

        _gemm_a16wfp4_preshuffle_kernel[p["grid"]](
            A,
            B_k,
            C_out,
            B_s,
            M,
            N,
            p["K_int"],
            A.stride(0),
            A.stride(1),
            B_k.stride(0),
            B_k.stride(1),
            0 if p["num_ksplit"] == 1 else y_pp.stride(0),
            C_out.stride(0) if p["num_ksplit"] == 1 else y_pp.stride(1),
            C_out.stride(1) if p["num_ksplit"] == 1 else y_pp.stride(2),
            B_s.stride(0),
            B_s.stride(1),
            PREQUANT=True,
            BLOCK_SIZE_M=p["bm"],
            BLOCK_SIZE_N=p["bn"],
            BLOCK_SIZE_K=p["bk"],
            GROUP_SIZE_M=1,
            NUM_KSPLIT=p["num_ksplit"],
            SPLITK_BLOCK_SIZE=p["spk_sz"],
            num_warps=p["num_warps"],
            num_stages=p.get("num_stages", 1),
            waves_per_eu=p["waves_per_eu"],
            matrix_instr_nonkdim=16,
            cache_modifier=p["cache_modifier"],
        )

        if main_keys_before is not None:
            r = _extract_runner(main_kc, main_keys_before, p["grid"])
            if r is not None:
                _RUNNER[shape_key] = r
                _RUNNER_ARGS[shape_key] = _build_main_args(p, A, B_k, C_out, B_s)
        ra = _RUNNER_ARGS.get(shape_key)
        if ra is not None:
            ra[1] = B_k
            ra[3] = B_s

        if p["num_ksplit"] > 1:
            red_kc, _ = _get_kernel_cache(_gemm_afp4wfp4_reduce_kernel)
            red_keys_before = set(red_kc.keys()) if red_kc is not None else None

            _gemm_afp4wfp4_reduce_kernel[p["reduce_grid"]](
                y_pp,
                C,
                M,
                N,
                y_pp.stride(0),
                y_pp.stride(1),
                y_pp.stride(2),
                C.stride(0),
                C.stride(1),
                16,
                16,
                p["actual_ksplit"],
                p["max_ksplit"],
            )

            if red_keys_before is not None:
                rr = _extract_runner(red_kc, red_keys_before, p["reduce_grid"])
                if rr is not None:
                    _REDUCE_RUNNER[shape_key] = rr
                    _REDUCE_RUNNER_ARGS[shape_key] = _build_reduce_args(p, y_pp, C)

        return C

    ra = _RUNNER_ARGS[shape_key]
    ra[0] = A
    ra[1] = B_k  # B changes each seed - always refresh
    ra[3] = B_s  # B changes each seed - always refresh
    runner(*ra)

    if p["num_ksplit"] > 1:
        if reduce_runner is None:
            _gemm_afp4wfp4_reduce_kernel[p["reduce_grid"]](
                y_pp,
                C,
                M,
                N,
                y_pp.stride(0),
                y_pp.stride(1),
                y_pp.stride(2),
                C.stride(0),
                C.stride(1),
                16,
                16,
                p["actual_ksplit"],
                p["max_ksplit"],
            )
            red_kc, _ = _get_kernel_cache(_gemm_afp4wfp4_reduce_kernel)
            if red_kc is not None:
                rr = _extract_runner(red_kc, set(), p["reduce_grid"])
                if rr is not None:
                    _REDUCE_RUNNER[shape_key] = rr
                    _REDUCE_RUNNER_ARGS[shape_key] = _build_reduce_args(p, y_pp, C)
        else:
            reduce_runner(*_REDUCE_RUNNER_ARGS[shape_key])

    return C


def _custom_kernel_quant_asm(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
    M: int,
    N: int,
    K_bf16: int,
) -> torch.Tensor:
    shape_key = (M, N, K_bf16)

    state = _STATES.get(shape_key)
    if state is None:
        state = _init_state(M, N, K_bf16, A.device)
        _STATES[shape_key] = state

    A_q_buf = state["A_q_buf"]
    A_scale_sh_buf = state["A_scale_sh_buf"]
    out_buf = state["out_buf"]
    A_q_fp4 = state["A_q_fp4"]
    A_scale_fp8 = state["A_scale_fp8"]
    cfg = state["cfg"]
    quant_runner = state["quant_runner"]

    if quant_runner is None:
        kc, _ = _get_kernel_cache(_mxfp4_quant_shuffled_kernel)
        kb = set(kc.keys()) if kc is not None else None

        _mxfp4_quant_shuffled_kernel[cfg["grid"]](
            A,
            A_q_buf,
            A_scale_sh_buf,
            A.stride(0),
            A.stride(1),
            A_q_buf.stride(0),
            A_q_buf.stride(1),
            M,
            K_bf16,
            state["SN"],
            BLOCK_SIZE_M=cfg["BLOCK_M"],
            BLOCK_SIZE_N=cfg["BLOCK_N"],
            NUM_ITER=cfg["NUM_ITER"],
            NUM_STAGES=cfg["NUM_STAGES"],
            MXFP4_QUANT_BLOCK_SIZE=32,
            EVEN_M_N=cfg["EVEN_M_N"],
            num_warps=cfg["NUM_WARPS"],
            waves_per_eu=2,
            num_stages=1,
        )

        if kb is not None:
            r = _extract_runner(kc, kb, cfg["grid"])
            if r is not None:
                state["quant_runner"] = r
                state["quant_ra"] = _build_quant_args(
                    cfg,
                    state["SN"],
                    A,
                    A_q_buf,
                    A_scale_sh_buf,
                    M,
                    K_bf16,
                )
    else:
        state["quant_ra"][0] = A
        quant_runner(*state["quant_ra"])

    gemm_a4w4_asm(
        A_q_fp4,
        B_shuffle,
        A_scale_fp8,
        B_scale_sh,
        out_buf,
        state["gemm_kernel_name"],
        None,
        1.0,
        0.0,
        True,
        state.get("log2_k_split", 0),
    )

    return out_buf[:M]


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

    M, K_bf16 = A.shape
    N = B_shuffle.shape[0]
    shape_key = (M, N, K_bf16)

    dispatch = _DISPATCH_CACHE.get(shape_key)
    if dispatch is not None:
        return dispatch(A, B_shuffle, B_scale_sh)

    fn = (
        _custom_kernel_preshuffle
        if _use_preshuffle(M, N, K_bf16)
        else _custom_kernel_quant_asm
    )
    result = fn(A, B_shuffle, B_scale_sh, M, N, K_bf16)

    def _bound(_A, _b, _bs, _fn=fn, _M=M, _N=N, _K=K_bf16):
        return _fn(_A, _b, _bs, _M, _N, _K)

    _DISPATCH_CACHE[shape_key] = _bound
    return result
scrolls · 790 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