Skip to content
KernelIndex
Search⌘K

submission 475463

Jitesh Jain US · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-475463?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
55.4µs
#207 of 310
2026-02-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a0ab68618bb4b2ae08fecdeb8b12cc23db7967b9beee8f371852ac065caa10e2
license declaredunknown
license concludedunknown
authorsJitesh Jain US
imported2026-08-15

Kernel source

submission.py203 lines
"""
Optimization strategy for v8:
- Use CUTLASS SM100 grouped GEMM extension with 1SM/2SM kernels
- Dispatch 2SM for large tiles, 1SM for small tiles
- Keep cached pointer arrays to avoid per-call overhead
"""

from __future__ import annotations

import os
from typing import Any

import torch

from task import input_t, output_t


def _flatten_reordered_scale(scale_reordered: torch.Tensor, l_idx: int) -> torch.Tensor:
    if scale_reordered.device.type != "cuda":
        scale_reordered = scale_reordered.cuda()

    scale_l = scale_reordered[..., l_idx]
    scale_l = scale_l.permute(2, 4, 0, 1, 3).contiguous()
    return scale_l.view(-1, 32, 16).flatten()


def _fallback_scaled_mm(
    abc_tensors,
    sfasfb_reordered_tensors,
    problem_sizes,
):
    result_tensors: list[torch.Tensor] = []
    for (a_ref, b_ref, c_ref), (sfa_reordered, sfb_reordered), (_m, _n, _k, l) in zip(
        abc_tensors, sfasfb_reordered_tensors, problem_sizes
    ):
        for l_idx in range(l):
            scale_a = _flatten_reordered_scale(sfa_reordered, l_idx)
            scale_b = _flatten_reordered_scale(sfb_reordered, l_idx)
            res = torch._scaled_mm(
                a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2),
                b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2),
                scale_a,
                scale_b,
                bias=None,
                out_dtype=torch.float16,
            )
            c_ref[:, :, l_idx] = res
        result_tensors.append(c_ref)
    return result_tensors


_EXT_MODULE = None
_PLAN_CACHE: dict[Any, dict[str, int]] = {}
_PTR_CACHE: dict[Any, dict[str, Any]] = {}


def _load_ext() -> Any:
    global _EXT_MODULE
    if _EXT_MODULE is not None:
        return _EXT_MODULE

    from torch.utils.cpp_extension import load

    this_dir = os.path.dirname(__file__)
    sources = [os.path.join(this_dir, "cutlass_grouped_gemm_ext.cu")]
    build_dir = os.path.join(this_dir, ".cutlass_build")
    os.makedirs(build_dir, exist_ok=True)

    cutlass_root = os.path.abspath(os.path.join(this_dir, "..", "..", "cutlass"))
    include_dirs = [os.path.join(cutlass_root, "include")]

    extra_cuda_cflags = [
        "-O3",
        "--use_fast_math",
        "-lineinfo",
        "-std=c++17",
        "-gencode=arch=compute_100,code=sm_100",
        "-DCUTLASS_ARCH_MMA_SM100_SUPPORTED",
    ]
    extra_cflags = [
        "-O3",
        "-std=c++17",
        "-DCUTLASS_ARCH_MMA_SM100_SUPPORTED",
    ]

    _EXT_MODULE = load(
        name="cutlass_grouped_gemm_ext",
        sources=sources,
        extra_include_paths=include_dirs,
        extra_cuda_cflags=extra_cuda_cflags,
        extra_cflags=extra_cflags,
        with_cuda=True,
        verbose=False,
        build_directory=build_dir,
    )
    return _EXT_MODULE


def _get_plans(problem_sizes):
    key = tuple(problem_sizes)
    cached = _PLAN_CACHE.get(key)
    if cached is not None:
        return cached

    ext = _load_ext()
    problem_sizes_mnk = torch.tensor(
        [(m, n, k) for (m, n, k, _l) in problem_sizes], dtype=torch.int32
    )
    handle_1 = ext.create_plan_1sm(problem_sizes_mnk)
    handle_2 = ext.create_plan_2sm(problem_sizes_mnk)
    _PLAN_CACHE[key] = {"1sm": handle_1, "2sm": handle_2}
    return _PLAN_CACHE[key]


def _get_ptr_bundle(key):
    return _PTR_CACHE.get(key)


def _set_ptr_bundle(key, bundle):
    _PTR_CACHE[key] = bundle


def _select_kernel(problem_sizes):
    # Prefer 2SM when N is large and M is at least 128.
    max_m = max(m for (m, _n, _k, _l) in problem_sizes)
    max_n = max(n for (_m, n, _k, _l) in problem_sizes)
    if max_m >= 128 and max_n >= 4096:
        return "2sm"
    return "1sm"


def custom_kernel(data: input_t) -> output_t:
    abc_tensors, _sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data

    try:
        handles = _get_plans(problem_sizes)
        kernel = _select_kernel(problem_sizes)
        handle = handles[kernel]

        key = (kernel,) + tuple(
            (a.data_ptr(), b.data_ptr(), c.data_ptr(), sfa.data_ptr(), sfb.data_ptr())
            for (a, b, c), (sfa, sfb) in zip(abc_tensors, sfasfb_reordered_tensors)
        )

        bundle = _get_ptr_bundle(key)
        if bundle is None:
            ptrs_a = []
            ptrs_b = []
            ptrs_c = []
            ptrs_sfa = []
            ptrs_sfb = []

            outputs = []
            for (a_ref, b_ref, c_ref), (sfa_reordered, sfb_reordered), (m, n, _k, l) in zip(
                abc_tensors, sfasfb_reordered_tensors, problem_sizes
            ):
                a_mat = a_ref[:, :, 0].contiguous()
                b_mat = b_ref[:, :, 0].contiguous()
                c_out = c_ref
                sfa_mat = sfa_reordered[:, :, :, :, :, 0].contiguous()
                sfb_mat = sfb_reordered[:, :, :, :, :, 0].contiguous()

                ptrs_a.append(a_mat.data_ptr())
                ptrs_b.append(b_mat.data_ptr())
                ptrs_c.append(c_out[:, :, 0].data_ptr())
                ptrs_sfa.append(sfa_mat.data_ptr())
                ptrs_sfb.append(sfb_mat.data_ptr())
                outputs.append(c_out)

            bundle = {
                "ptrs_a": torch.tensor(ptrs_a, device="cuda", dtype=torch.int64),
                "ptrs_b": torch.tensor(ptrs_b, device="cuda", dtype=torch.int64),
                "ptrs_c": torch.tensor(ptrs_c, device="cuda", dtype=torch.int64),
                "ptrs_sfa": torch.tensor(ptrs_sfa, device="cuda", dtype=torch.int64),
                "ptrs_sfb": torch.tensor(ptrs_sfb, device="cuda", dtype=torch.int64),
                "outputs": outputs,
            }
            _set_ptr_bundle(key, bundle)

        ext = _load_ext()
        if kernel == "2sm":
            ext.run_plan_2sm(
                handle,
                bundle["ptrs_a"],
                bundle["ptrs_b"],
                bundle["ptrs_c"],
                bundle["ptrs_sfa"],
                bundle["ptrs_sfb"],
            )
        else:
            ext.run_plan_1sm(
                handle,
                bundle["ptrs_a"],
                bundle["ptrs_b"],
                bundle["ptrs_c"],
                bundle["ptrs_sfa"],
                bundle["ptrs_sfb"],
            )
        torch.cuda.synchronize()
        return bundle["outputs"]
    except Exception:
        return _fallback_scaled_mm(abc_tensors, sfasfb_reordered_tensors, problem_sizes)
scrolls · 203 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