Skip to content
KernelIndex
Search⌘K

submission 596893

inference_and_chill · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:69fd29ff4adc362b029aaa783b31e2e12cdd8c141d7a79b857e849d863639c41
license declaredunknown
license concludedunknown
authorsinference_and_chill
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM with per-input dispatch table.

Kernel source

submission.py100 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
FP4 quant + FP4 GEMM with per-input dispatch table.

Benchmark cases (from task.yml): (m, n, k)
  (4,   2880,  512)
  (16,  2112, 7168)
  (32,  4096,  512)
  (32,  2880,  512)
  (64,  7168, 2048)
  (256, 3072, 1536)

PER_CASE_CONFIGS maps (m, n, k) -> config dict.
All entries start empty (falling back to DEFAULT_CONFIG).
Fill in per-case params after benchmarking.
"""

from typing import TYPE_CHECKING, Any

if TYPE_CHECKING:
    from task import input_t, output_t
else:
    input_t = Any
    output_t = Any

try:
    # Triton path — may not be available on all builds
    # Actual module: aiter.ops.triton.gemm.basic.gemm_afp4wfp4
    # NOTE: the import path may differ across aiter versions; adjust if needed.
    from aiter.ops.triton.gemm.basic import gemm_afp4wfp4 as triton_gemm_a4w4  # noqa: F401
    _HAS_TRITON_GEMM = True
except (ImportError, ModuleNotFoundError):
    triton_gemm_a4w4 = None
    _HAS_TRITON_GEMM = False

# ---------------------------------------------------------------------------
# Per-case dispatch table
# key: (m, n, k)  — A is (m, k), B is (n, k)
# value: dict of kernel params; empty dict → use DEFAULT_CONFIG
# ---------------------------------------------------------------------------
PER_CASE_CONFIGS: dict[tuple[int, int, int], dict] = {
    (4,   2880,  512): {},  # TODO: tune
    (16,  2112, 7168): {},  # TODO: tune
    (32,  4096,  512): {},  # TODO: tune
    (32,  2880,  512): {},  # TODO: tune
    (64,  7168, 2048): {},  # TODO: tune
    (256, 3072, 1536): {},  # TODO: tune
}

DEFAULT_CONFIG: dict = {
    "engine": "ck",  # "ck" | "triton"
}


def _run_kernel(data, cfg: dict):
    """Execute gemm_a4w4 using config-specified engine."""
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    engine = cfg.get("engine", "ck")

    def _quant_mxfp4(x, shuffle=True):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

    A, B, _B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()

    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)

    if engine == "triton" and _HAS_TRITON_GEMM:
        return triton_gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh)

    # Default: CK path
    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def custom_kernel(data: input_t) -> output_t:
    A, B, _B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape   # A: (m, k)
    n, _ = B.shape   # B: (n, k)

    key = (int(m), int(n), int(k))
    case_cfg = PER_CASE_CONFIGS.get(key, {})
    cfg = {**DEFAULT_CONFIG, **case_cfg}
    return _run_kernel(data, cfg)
scrolls · 100 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