Skip to content
KernelIndex
Search⌘K

submission 748568

coderwhisper · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v104_fixptr2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-748568?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
10.0µs
#236 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:521e61c3f675a087d0ca3e68a12de2d1b96d84dbdf1c1e177bbea46e650f1f7e
license declaredunknown
license concludedunknown
authorscoderwhisper
imported2026-08-15

Techniques

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

split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
stages = 1num_warps=q_NW, waves_per_eu=0, num_stages=1,

Kernel source

submission_v104_fixptr2.py638 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
V104: Fixed direct launcher.launch bypass (attempt 2).
V103's fix failed because captured args contain tensor OBJECTS, not int pointers.
isinstance(val, int) matched nothing -> ptr_subs=[0] for all shapes.

Fix: match both torch.Tensor (by data_ptr()) and int args. For tensor args,
substitute with a tensor sharing the new data's storage. For int args, substitute
with new data_ptr() values.
"""
import torch
import triton
import triton.language as tl
import sys
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_preshuffle_kernel,
)
from task import input_t, output_t

SCALE_GROUP = 32
_cache = {}
_bf16 = torch.bfloat16

def p(*args):
    print(*args, file=sys.stderr)


# ===================== A16WFP4 preshuffle kernel with inline KSPLIT reduction =====================

@triton.heuristics(
    {
        "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
    }
)
@triton.jit
def _a16wfp4_inline_reduce_kernel(
    a_ptr, b_ptr, c_ptr, c_final_ptr, b_scales_ptr, counter_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
    stride_cf_m, stride_cf_n,
    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,
    EVEN_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,
    INLINE_REDUCE: tl.constexpr,
):
    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)

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    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)
    tl.assume(pid_k >= 0)

    SCALE_GROUP_SIZE: tl.constexpr = 32

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)

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

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

        offs_bsn = (
            pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
        ) % N
        offs_ks = (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[None, :] * stride_bsk
        )

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

        for k in tl.range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter, num_stages=num_stages):
            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)
            )

            if EVEN_K:
                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(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

        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)

        if INLINE_REDUCE:
            c_ptrs = (
                c_ptr
                + pid_k * stride_ck
                + stride_cm * offs_cm[:, None]
                + stride_cn * offs_cn[None, :]
            )
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(c_ptrs, accumulator, mask=c_mask, cache_modifier=".wt")

            tile_id = pid_m * num_pid_n + pid_n
            old_count = tl.atomic_add(counter_ptr + tile_id, 1)

            if old_count == NUM_KSPLIT - 1:
                total = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
                for ks in range(NUM_KSPLIT):
                    partial_ptrs = (
                        c_ptr
                        + ks * stride_ck
                        + stride_cm * offs_cm[:, None]
                        + stride_cn * offs_cn[None, :]
                    )
                    partial = tl.load(partial_ptrs, mask=c_mask)
                    total += partial

                result = total.to(c_final_ptr.type.element_ty)
                cf_ptrs = (
                    c_final_ptr
                    + stride_cf_m * offs_cm[:, None]
                    + stride_cf_n * offs_cn[None, :]
                )
                tl.store(cf_ptrs, result, mask=c_mask, cache_modifier=".wt")

                tl.atomic_xchg(counter_ptr + tile_id, 0)
        else:
            c = accumulator.to(c_final_ptr.type.element_ty)
            cf_ptrs = (
                c_final_ptr
                + stride_cf_m * offs_cm[:, None]
                + stride_cf_n * offs_cn[None, :]
            )
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(cf_ptrs, c, mask=c_mask, cache_modifier=".wt")


# ===================== Fused quant+shuffle kernel =====================

@triton.jit
def _fused_quant_shuffle_kernel(
    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, padded_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,
    SCALING_MODE: 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_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(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)
        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)
        i = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        j = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        shuffled_off = (
            (i // 32)[:, None] * (32 * padded_sn)
            + (j // 8)[None, :] * 256 + (j % 4)[None, :] * 64
            + (i % 16)[:, None] * 4 + ((j % 8) // 4)[None, :] * 2
            + ((i % 32) // 16)[:, None]
        )
        if EVEN_M_N:
            tl.store(bs_ptr + shuffled_off, bs_e8m0)
        else:
            scale_n = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            bs_mask = (i < M)[:, None] & (j < scale_n)[None, :]
            tl.store(bs_ptr + shuffled_off, bs_e8m0, mask=bs_mask)


# ===================== Split-K helper =====================

def _get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):
    SPLITK_BLOCK_SIZE = (
        triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
    )
    while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
        if (K % (SPLITK_BLOCK_SIZE // 2) == 0
            and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
            and K % (BLOCK_SIZE_K // 2) == 0):
            break
        elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
            NUM_KSPLIT = NUM_KSPLIT // 2
        elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
            if NUM_KSPLIT > 1:
                NUM_KSPLIT = NUM_KSPLIT // 2
            elif BLOCK_SIZE_K > 16:
                BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
            BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        else:
            break
        SPLITK_BLOCK_SIZE = (
            triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
        )
    return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT


# ===================== Per-shape configs =====================

_PS_CONFIGS = {
    "k_small_m4": {
        "BLOCK_SIZE_M": 4, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2,
        "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    "k_small_m32": {
        "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    "k7168": {
        "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 7,
    },
}

_AFP4_CONFIGS = {
    "k2048_m64": {
        "BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 1024,
        "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    "k1536_m256": {
        "BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": None, "NUM_KSPLIT": 1,
    },
}


# ===================== A16WFP4 Preshuffle path =====================

def _prepare_preshuffle(M, K, N, dev):
    if K <= 1024:
        cfg = dict(_PS_CONFIGS["k_small_m4" if M <= 4 else "k_small_m32"])
    else:
        cfg = dict(_PS_CONFIGS["k7168"])

    K_packed = K // 2
    if cfg["NUM_KSPLIT"] > 1:
        sbs, bsk, nks = _get_splitk(K_packed, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])
        cfg["SPLITK_BLOCK_SIZE"] = sbs
        cfg["BLOCK_SIZE_K"] = bsk
        cfg["NUM_KSPLIT"] = nks
    if cfg["BLOCK_SIZE_K"] >= 2 * K_packed:
        cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
        cfg["NUM_KSPLIT"] = 1
    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)

    use_splitk = cfg["NUM_KSPLIT"] > 1
    y = torch.empty((M, N), dtype=_bf16, device=dev)

    if use_splitk:
        y_pp = torch.empty((cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=dev)
        num_pid_m = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
        num_pid_n = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
        counter = torch.zeros(num_pid_m * num_pid_n, dtype=torch.int32, device=dev)
    else:
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
        y_pp = None
        counter = None

    grid = (cfg["NUM_KSPLIT"] * triton.cdiv(M, cfg["BLOCK_SIZE_M"]) * triton.cdiv(N, cfg["BLOCK_SIZE_N"]),)
    return ("ps", cfg, grid, y, y_pp, use_splitk, K_packed, counter)


def _run_preshuffle(data, M, K, N, cached):
    _, cfg, grid, y, y_pp, use_splitk, K_packed, counter = cached

    w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
    bs_uint8 = data[4].view(torch.uint8)
    sm, sn = bs_uint8.shape
    w_scales = bs_uint8.reshape(sm // 32, sn * 32)

    if use_splitk:
        _a16wfp4_inline_reduce_kernel[grid](
            data[0], w_ps, y_pp, y, w_scales, counter,
            M, N, K_packed,
            data[0].stride(0), data[0].stride(1),
            w_ps.stride(0), w_ps.stride(1),
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            w_scales.stride(0), w_scales.stride(1),
            y.stride(0), y.stride(1),
            INLINE_REDUCE=True,
            **cfg,
        )
    else:
        _a16wfp4_inline_reduce_kernel[grid](
            data[0], w_ps, None, y, w_scales, None,
            M, N, K_packed,
            data[0].stride(0), data[0].stride(1),
            w_ps.stride(0), w_ps.stride(1),
            0, y.stride(0), y.stride(1),
            w_scales.stride(0), w_scales.stride(1),
            y.stride(0), y.stride(1),
            INLINE_REDUCE=False,
            **cfg,
        )
    return y


# ===================== AFP4WFP4 preshuffle path (medium K) =====================

def _prepare_afp4(M, K, N, dev):
    K_packed = K // 2
    x_fp4 = torch.empty((M, K_packed), dtype=torch.uint8, device=dev)
    sn = triton.cdiv(K, SCALE_GROUP)
    padded_sm = triton.cdiv(M, 256) * 256
    padded_sn = triton.cdiv(sn, 8) * 8
    bs_shuffled = torch.zeros((padded_sm, padded_sn), dtype=torch.uint8, device=dev)

    if M <= 32:
        q_NI, q_BSM, q_BSN, q_NW, q_NS = 1, triton.next_power_of_2(M), 32, 1, 1
    else:
        q_NI, q_BSM, q_BSN, q_NW, q_NS = 4, 32, 128, 4, 2
    q_EVEN = (M % q_BSM == 0) and (K % q_BSN == 0)
    grid_q = (triton.cdiv(M, q_BSM), triton.cdiv(K, q_BSN * q_NI))
    fp4_s = x_fp4.stride()

    if K == 2048 and M <= 64:
        cfg = dict(_AFP4_CONFIGS["k2048_m64"])
    else:
        cfg = dict(_AFP4_CONFIGS["k1536_m256"])

    if cfg["NUM_KSPLIT"] > 1:
        sbs, bsk, nks = _get_splitk(K_packed, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])
        cfg["SPLITK_BLOCK_SIZE"] = sbs
        cfg["BLOCK_SIZE_K"] = bsk
        cfg["NUM_KSPLIT"] = nks
    else:
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed

    if cfg["BLOCK_SIZE_K"] >= 2 * K_packed:
        cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)

    grid_gemm = (
        cfg["NUM_KSPLIT"]
        * triton.cdiv(M, cfg["BLOCK_SIZE_M"])
        * triton.cdiv(N, cfg["BLOCK_SIZE_N"]),
    )

    y = torch.empty((M, N), dtype=_bf16, device=dev)

    return ("afp4", x_fp4, bs_shuffled, padded_sm, padded_sn,
            cfg, y, grid_q, q_BSM, q_BSN, q_NI, q_NS, q_NW, q_EVEN,
            fp4_s[0], fp4_s[1], grid_gemm, K_packed)


def _run_afp4(data, M, K, N, cached):
    (_, x_fp4, bs_sh, padded_sm, padded_sn,
     cfg, y, grid_q, q_BSM, q_BSN, q_NI, q_NS, q_NW, q_EVEN,
     fp4s0, fp4s1, grid_gemm, K_packed) = cached

    A = data[0]

    _fused_quant_shuffle_kernel[grid_q](
        A, x_fp4, bs_sh,
        A.stride(0), A.stride(1), fp4s0, fp4s1,
        M, K, padded_sn,
        BLOCK_SIZE_M=q_BSM, BLOCK_SIZE_N=q_BSN,
        NUM_ITER=q_NI, NUM_STAGES=q_NS,
        MXFP4_QUANT_BLOCK_SIZE=SCALE_GROUP,
        SCALING_MODE=0, EVEN_M_N=q_EVEN,
        num_warps=q_NW, waves_per_eu=0, num_stages=1,
    )

    a_scales = bs_sh.reshape(padded_sm // 32, padded_sn * 32)
    w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
    bs_uint8 = data[4].view(torch.uint8)
    bsm, bsn = bs_uint8.shape
    w_scales = bs_uint8.reshape(bsm // 32, bsn * 32)

    _gemm_afp4wfp4_preshuffle_kernel[grid_gemm](
        x_fp4, w_ps, y, a_scales, w_scales,
        M, N, K_packed,
        x_fp4.stride(0), x_fp4.stride(1),
        w_ps.stride(0), w_ps.stride(1),
        0, y.stride(0), y.stride(1),
        a_scales.stride(0), a_scales.stride(1),
        w_scales.stride(0), w_scales.stride(1),
        **cfg,
    )

    return y


# ===================== Direct launch bypass infrastructure =====================

# key -> list of (launch_fn, args_list, ptr_subs)
# ptr_subs: list of (position, data_index, is_tensor)
_fast_cache = {}
_outputs = {}      # key -> output tensor


def _get_jit_fn(kfn):
    while not hasattr(kfn, 'device_caches') and hasattr(kfn, 'fn'):
        kfn = kfn.fn
    return kfn


def _patch_and_capture(kernel_fns, run_fn):
    captured = []
    patches = []

    for kfn in kernel_fns:
        jit_fn = _get_jit_fn(kfn)
        if not hasattr(jit_fn, 'device_caches') or 0 not in jit_fn.device_caches:
            continue
        cache_dict = jit_fn.device_caches[0][0]
        for k, ck in cache_dict.items():
            try:
                orig = ck.run.launch
                def make_hook(o):
                    def hook(*a):
                        captured.append((o, a))
                        return o(*a)
                    return hook
                new_fn = make_hook(orig)
                ck.run.launch = new_fn
                if ck.run.launch is new_fn:
                    patches.append((ck.run, orig))
                else:
                    ck.run.launch = orig
            except (AttributeError, TypeError):
                pass

    result = run_fn()

    for launcher_obj, orig in patches:
        try:
            launcher_obj.launch = orig
        except (AttributeError, TypeError):
            pass

    return captured, result


# ===================== Derive input tensors for fast path =====================

def _derive_tensors_ps(data, K_packed, N):
    """Derive kernel arg tensors for preshuffle path."""
    w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
    bs_uint8 = data[4].view(torch.uint8)
    sm, sn = bs_uint8.shape
    w_scales = bs_uint8.reshape(sm // 32, sn * 32)
    return data[0], w_ps, w_scales


def _derive_tensors_afp4(data, K_packed, N):
    """Derive kernel arg tensors for AFP4WFP4 path."""
    w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
    bs_uint8 = data[4].view(torch.uint8)
    bsm, bsn = bs_uint8.shape
    w_scales = bs_uint8.reshape(bsm // 32, bsn * 32)
    return data[0], w_ps, w_scales


# ===================== Router =====================

def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    M, K = A.shape
    N = data[3].shape[0]
    key = (M, K, N)

    # Fast path: replay with updated input tensor pointers
    if key in _fast_cache:
        fc = _fast_cache[key]
        if fc is not None:
            entries, derive_fn, K_packed = fc
            a_new, w_ps_new, w_scales_new = derive_fn(data, K_packed, N)
            new_tensors = {0: a_new, 3: w_ps_new, 4: w_scales_new}
            for launch_fn, args_list, ptr_subs in entries:
                for pos, data_idx, is_tensor in ptr_subs:
                    if is_tensor:
                        args_list[pos] = new_tensors[data_idx]
                    else:
                        args_list[pos] = new_tensors[data_idx].data_ptr()
                launch_fn(*args_list)
            return _outputs[key]
        else:
            cached = _cache[key]
            if cached[0] == "ps":
                return _run_preshuffle(data, M, K, N, cached)
            else:
                return _run_afp4(data, M, K, N, cached)

    # First call: prepare, compile, run, and capture
    if key not in _cache:
        if K <= 1024 or K >= 4096:
            _cache[key] = _prepare_preshuffle(M, K, N, A.device)
        else:
            _cache[key] = _prepare_afp4(M, K, N, A.device)

    cached = _cache[key]

    # Run once to compile
    if cached[0] == "ps":
        result = _run_preshuffle(data, M, K, N, cached)
        kfns = [_a16wfp4_inline_reduce_kernel]
        run_fn = lambda: _run_preshuffle(data, M, K, N, cached)
        derive_fn = _derive_tensors_ps
        K_packed = K // 2
    else:
        result = _run_afp4(data, M, K, N, cached)
        kfns = [_fused_quant_shuffle_kernel, _gemm_afp4wfp4_preshuffle_kernel]
        run_fn = lambda: _run_afp4(data, M, K, N, cached)
        derive_fn = _derive_tensors_afp4
        K_packed = K // 2

    # Record input tensor pointers for substitution mapping
    # data[0]=A, data[3]=B_shuffle (views share same data_ptr), data[4]=B_scale_sh
    input_ptrs = {}
    for idx in (0, 3, 4):
        ptr = data[idx].data_ptr()
        if ptr not in input_ptrs:
            input_ptrs[ptr] = idx

    # Capture launcher.launch args by running again with patches
    captures, _ = _patch_and_capture(kfns, run_fn)

    if captures:
        fast_entries = []
        for fn, args in captures:
            args_list = list(args)
            ptr_subs = []
            for pos, val in enumerate(args_list):
                # Check int pointers
                if isinstance(val, int) and val in input_ptrs:
                    ptr_subs.append((pos, input_ptrs[val], False))
                # Check tensor objects (Triton may pass tensors, not ints)
                elif isinstance(val, torch.Tensor):
                    try:
                        dptr = val.data_ptr()
                        if dptr in input_ptrs:
                            ptr_subs.append((pos, input_ptrs[dptr], True))
                    except Exception:
                        pass
            fast_entries.append((fn, args_list, ptr_subs))

        _fast_cache[key] = (fast_entries, derive_fn, K_packed)
        _outputs[key] = result
        p(f"[V104] {key}: captured {len(captures)} launches, ptr_subs={[len(e[2]) for e in fast_entries]}")
    else:
        _fast_cache[key] = None
        p(f"[V104] {key}: capture FAILED, using normal path")

    return result
scrolls · 638 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