Skip to content
KernelIndex
Search⌘K

submission 80635

akzaidan · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

cute2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-80635?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 GEMVsuite of 3 cases
NVIDIA B200
43.9µs
#247 of 678
2025-11-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0f9e350c088e8029df6cc43caa82087d2cb42ac356f6d8bc67b2df1406765508
license declaredunknown
license concludedunknown
authorsakzaidan
imported2026-08-15

Kernel source

cute2.py230 lines
from task import input_t, output_t

import cutlass  # pyright: ignore[reportMissingImports]
import cutlass.cute as cute  # pyright: ignore[reportMissingImports]
from cutlass.cute.runtime import make_ptr  # pyright: ignore[reportMissingImports]
import cutlass.utils.blockscaled_layout as blockscaled_utils  # pyright: ignore[reportMissingImports]

# ---------------------------------------------------------------------------
# Kernel Config
# ---------------------------------------------------------------------------
mma_tiler_mnk = (128, 1, 64)
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
threads_per_cta = 128
sf_vec_size = 16


# ---------------------------------------------------------------------------
# Helper: ceil div
# ---------------------------------------------------------------------------
def ceil_div(a, b):
    return (a + b - 1) // b


# ---------------------------------------------------------------------------
# FP32 atomic add wrapper (required for parallel-K)
# ---------------------------------------------------------------------------
from cutlass import Float32
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm, llvm


@dsl_user_op
def atomic_add_fp32(a: float | Float32, gmem_ptr: cute.Pointer, *, loc=None, ip=None) -> None:
    nvvm.atomicrmw(
        res=T.f32(),
        op=nvvm.AtomicOpKind.FADD,
        ptr=gmem_ptr.llvm_ptr,
        a=Float32(a).ir_value(),
    )


# ---------------------------------------------------------------------------
# PARALLEL-K KERNEL (each block handles exactly one K tile)
# ---------------------------------------------------------------------------
@cute.kernel
def kernel(
    mA_mkl: cute.Tensor,
    mB_nkl: cute.Tensor,
    mSFA_mkl: cute.Tensor,
    mSFB_nkl: cute.Tensor,
    mCaccum_mnl: cute.Tensor,   # IMPORTANT: this is FLOAT32 accumulation buffer
):
    bidx, bidy, bidz = cute.arch.block_idx()  # (M tile, K tile, L batch)
    tidx, _, _ = cute.arch.thread_idx()

    # ---- Local tiling (same as reference) ----
    gA_mkl = cute.local_tile(mA_mkl,  cute.slice_(mma_tiler_mnk, (None, 0, None)), (None,None,None))
    gSFA_mkl = cute.local_tile(mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None,None,None))

    gB_nkl = cute.local_tile(mB_nkl,  cute.slice_(mma_tiler_mnk, (0, None, None)), (None,None,None))
    gSFB_nkl = cute.local_tile(mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None,None,None))

    gC_mnl = cute.local_tile(mCaccum_mnl, cute.slice_(mma_tiler_mnk, (None,None,0)), (None,None,None))

    # Output slice - FP32 accumulation
    tCgC = gC_mnl[tidx, None, bidx, 0, bidz]
    tCgC = cute.make_tensor(tCgC.iterator, 1)

    # FP32 accumulator
    res = cute.zeros_like(tCgC, cutlass.Float32)

    # ------ SINGLE K-TILE (parallel K) ------
    k_tile = bidy

    tAgA   = gA_mkl[tidx, None, bidx, k_tile, bidz]
    tBgB   = gB_nkl[0, None, 0, k_tile, bidz]
    tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
    tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]

    # rmem staging
    tArA   = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
    tBrB   = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)
    tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
    tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)

    # load + convert
    a_val = tAgA.load().to(cutlass.Float16)
    b_val = tBgB.load().to(cutlass.Float16)
    sfa_val = tAgSFA.load().to(cutlass.Float32)
    sfb_val = tBgSFB.load().to(cutlass.Float32)

    # store to rmem
    tArA.store(a_val)
    tBrB.store(b_val)
    tArSFA.store(sfa_val)
    tBrSFB.store(sfb_val)

    # tilewise FFMA
    for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
        res += (tArA[i] * tBrB[i]) * (tArSFA[i] * tBrSFB[i])

    # accumulate
    atomic_add_fp32(res[0], tCgC.iterator)
    return


# ---------------------------------------------------------------------------
# JIT launcher (updated to compute parallel-K grid)
# ---------------------------------------------------------------------------
@cute.jit
def my_kernel(
    a_ptr: cute.Pointer,
    b_ptr: cute.Pointer,
    sfa_ptr: cute.Pointer,
    sfb_ptr: cute.Pointer,
    caccum_ptr: cute.Pointer,   # FP32 buffer, not FP16!
    problem_size: tuple,
):
    m, _, k, l = problem_size

    # A tensor (M x K x L)
    a_tensor = cute.make_tensor(
        a_ptr,
        cute.make_layout(
            (m, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(m*k, 32)),
        ),
    )

    # B tensor padded to n=128
    n_pad = 128
    b_tensor = cute.make_tensor(
        b_ptr,
        cute.make_layout(
            (n_pad, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(n_pad*k, 32)),
        ),
    )

    # FP32 accumulation buffer (same layout as C but FP32)
    c_tensor = cute.make_tensor(
        caccum_ptr,
        cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m)),
    )

    # Scale-factor layouts
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
    sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)

    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
    sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)

    # -----------------------------
    # PARALLEL-K grid:
    #   M_tiles = ceil(M/128)
    #   K_tiles = ceil(K/64)
    #   L       = batch
    # -----------------------------
    grid = (
        cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),  # M tiles
        cute.ceil_div(a_tensor.shape[1], mma_tiler_mnk[2]),  # K tiles
        c_tensor.shape[2],                                  # L batch
    )

    kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
        grid=grid,
        block=[threads_per_cta, 1, 1],
        cluster=(1, 1, 1),
    )
    return


# ---------------------------------------------------------------------------
# Compile the kernel once
# ---------------------------------------------------------------------------
_compiled_kernel_cache = None


def compile_kernel():
    global _compiled_kernel_cache
    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache

    a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    caccum_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
    sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)

    _compiled_kernel_cache = cute.compile(
        my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, caccum_ptr, (0, 0, 0, 0)
    )
    return _compiled_kernel_cache


# ---------------------------------------------------------------------------
# ENTRY POINT (calls parallel-K kernel)
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    a, b, _, _, sfa_perm, sfb_perm, c_fp16 = data

    compiled = compile_kernel()

    m, k_half, l = a.shape
    k = k_half * 2

    # -------------------------------
    # Allocate FP32 accumulation buffer
    # -------------------------------
    import torch

    c_accum = torch.zeros_like(c_fp16, dtype=torch.float32)

    a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    caccum_ptr = make_ptr(
        cutlass.Float32, c_accum.data_ptr(), cute.AddressSpace.gmem, assumed_align=16
    )

    sfa_ptr = make_ptr(sf_dtype, sfa_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
    sfb_ptr = make_ptr(sf_dtype, sfb_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)

    compiled(a_ptr, b_ptr, sfa_ptr, sfb_ptr, caccum_ptr, (m, 1, k, l))

    # Copy FP32 → FP16 user output
    c_fp16.copy_(c_accum.to(torch.float16))
    return c_fp16
scrolls · 230 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