Skip to content
KernelIndex
Search⌘K

submission 593457

abhicloudstalk13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

improved_gemm_submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-593457?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
18.6µs
#698 of 1143
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b2f657f665b39229340ccd7a05999d29d2df2f7166a02bfdeef41cc6b7556a44
license declaredunknown
license concludedunknown
authorsabhicloudstalk13
imported2026-08-26

Kernel source

improved_gemm_submission_v2.py92 lines
from task import input_t, output_t
import torch

from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle


def _c(bm, bn, bk, grp, nw, ns, wpe, ksplit, cm=None):
    return {
        "BLOCK_SIZE_M": bm,
        "BLOCK_SIZE_N": bn,
        "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": grp,
        "num_warps": nw,
        "num_stages": ns,
        "waves_per_eu": wpe,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": cm,
        "NUM_KSPLIT": ksplit,
    }


_CFG_2880_4 = _c(8, 32, 256, 1, 2, 2, 2, 2, ".cg")
_CFG_2880_32 = _c(16, 32, 256, 1, 4, 2, 2, 2, ".cg")
_CFG_4096_32 = _c(32, 32, 256, 1, 4, 2, 2, 2, ".cg")
_CFG_2112_16 = _c(16, 128, 512, 1, 4, 1, 1, 14, ".cg")
_CFG_7168_64 = _c(16, 64, 512, 1, 8, 2, 4, 1, ".cg")
_CFG_3072_256 = _c(64, 32, 512, 1, 4, 2, 1, 1)
_CFG_DEFAULT = _c(32, 64, 512, 1, 8, 1, 2, 1)

_WEIGHT_VIEW_CACHE = {}


def _prepare_weight_views(weight, scale, n, k):
    key = (
        weight.data_ptr(),
        scale.data_ptr(),
        weight.shape,
        scale.shape,
        n,
        k,
    )
    cached = _WEIGHT_VIEW_CACHE.get(key)
    if cached is None:
        packed_weight = weight.view(torch.uint8).view(n >> 4, (k >> 1) << 4)
        packed_scale = scale.view(torch.uint8).view(scale.shape[0] >> 5, k)[: n >> 5]
        cached = (packed_weight, packed_scale)
        _WEIGHT_VIEW_CACHE[key] = cached
    return cached


def _get_config(m, n, k):
    if k == 512:
        if n == 2880:
            if m <= 8:
                return _CFG_2880_4
            if m <= 32:
                return _CFG_2880_32
        elif n == 4096 and m <= 32:
            return _CFG_4096_32
    elif k == 7168:
        if n == 2112 and m <= 16:
            return _CFG_2112_16
    elif k == 2048:
        if n == 7168 and m <= 64:
            return _CFG_7168_64
    elif k == 1536:
        if n == 3072 and m <= 256:
            return _CFG_3072_256
    return _CFG_DEFAULT


def custom_kernel(data: input_t) -> output_t:
    a = data[0]
    b_shuffle = data[3]
    b_scale_sh = data[4]

    if a.stride(1) != 1:
        a = a.contiguous()

    m, k = a.shape
    n = b_shuffle.shape[0]
    b_w, b_scale_w = _prepare_weight_views(b_shuffle, b_scale_sh, n, k)

    return gemm_a16wfp4_preshuffle(
        a,
        b_w,
        b_scale_w,
        dtype=dtypes.bf16,
        config=_get_config(m, n, k),
    )
scrolls · 92 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