Skip to content
KernelIndex
Search⌘K

submission 86187

Emilio · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

kernel28.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-86187?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
32.3µs
#174 of 678
2025-11-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:68e11f18ca1c4927e731a982d42152d48bb0a25c9b845b760f4caa9ed06f73eb
license declaredunknown
license concludedunknown
authorsEmilio
imported2026-08-15

Techniques

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

fp4NVFP4 GEMV Optimization v28: 8-Way Unrolling
shared-memorysmem_layout = cute.make_layout(

Kernel source

kernel28.py317 lines
import torch
from task import input_t, output_t

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm

"""
NVFP4 GEMV Optimization v28: 8-Way Unrolling
===============================================

OPTIMIZATION: Increase unrolling factor from 4 to 8 for better ILP

Rationale:
- 4-way unrolling exposes some instruction-level parallelism
- 8-way unrolling can hide more latency (FMUL latency ~4 cycles on Blackwell)
- More independent operations allow better instruction scheduling
- Register pressure should still be manageable (K=128/256, unroll=8 → 16/32 iterations)

Expected improvements:
- 1.05-1.15× speedup from better instruction scheduling
- Lower impact than vectorization, but safer and simpler

Based on v26 (adaptive K-tiling baseline)
"""

sf_vec_size = 16
threads_per_m = 32
threads_per_k = 8
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16


def ceil_div(a, b):
    return (a + b - 1) // b


def get_k_tile_size(k):
    """
    Adaptive K-tile sizing based on K dimension:
    - Small K (≤2048): Use K=128 for better thread utilization
    - Large K (>2048): Use K=256 for cache-line alignment
    """
    if k <= 2048:
        return 128
    else:
        return 256


def get_threads_per_k(k):
    """
    Adaptive threads_per_k based on problem size:
    - Small K (≤2048): Use fewer threads (4) to reduce SMEM reduction overhead
    - Medium K (≤8192): Use 6 threads for balanced parallelism
    - Large K (>8192): Use 8 threads for more K-parallelism

    Empirical results:
    - k=2048: threads_per_k=4 → 17.4 µs (best)
    - k=7168: threads_per_k=6 → 50.0 µs (best)
    - k=16384: threads_per_k=8 → 35.6 µs (best)
    """
    if k <= 2048:
        return 4
    elif k <= 8192:
        return 6
    else:
        return 8


@dsl_user_op
def elem_pointer(x: cute.Tensor, coord: cute.Coord, *, loc=None, ip=None) -> cute.Pointer:
    """Get pointer to specific element in tensor for atomic operations."""
    return x.iterator + cute.crd2idx(coord, x.layout, loc=loc, ip=ip)


# Warp-reduction kernel with ADAPTIVE K-tiling and threads_per_k
@cute.kernel
def cutedsl_kernel(
    mA_mkl: cute.Tensor,
    mB_nkl: cute.Tensor,
    mSFA_mkl: cute.Tensor,
    mSFB_nkl: cute.Tensor,
    mC_mnl: cute.Tensor,
    k_tile_size: cutlass.Constexpr,
    threads_per_k_param: cutlass.Constexpr,
):
    bidx, bidy, bidz = cute.arch.block_idx()
    tidx, tidy, _ = cute.arch.thread_idx()

    # Create tile shape with dynamic K-tile size
    mma_tiler_mnk = (threads_per_m, 1, k_tile_size)

    # Extract tiles (using dynamic tile size)
    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(
        mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
    )

    # Output location (same for all tidy threads in this row)
    tCgC = gC_mnl[tidx, None, bidx, 0, bidz]
    tCgC = cute.make_tensor(tCgC.iterator, 1)
    res = cute.zeros_like(tCgC, cutlass.Float32)

    # CRITICAL: Allocate 2D SMEM with K-major layout
    allocator = cutlass.utils.SmemAllocator()
    smem_layout = cute.make_layout(
        (threads_per_m, threads_per_k_param),
        stride=(threads_per_k_param, 1)  # K-major: stride-1 in tidy dimension
    )
    shared_res = allocator.allocate_tensor(
        element_type=cutlass.Float32,
        layout=smem_layout
    )

    # Grid-stride loop over K-tiles
    # NOTE: k_tile_cnt now varies based on k_tile_size!
    # Example for case 2:
    #   v15: k=2048, K=256 → 8 tiles (0.5 tiles/thread)
    #   v16: k=2048, K=128 → 16 tiles (1.0 tiles/thread) ← BETTER!
    k_tile_cnt = gA_mkl.layout[3].shape

    for k_tile in range(tidy, k_tile_cnt, threads_per_k_param, unroll_full=True):
        # Load A and scale factors for this M-thread (tidx), this K-tile
        tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
        tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]

        # Load B and scale factors for this K-tile (shared across all M-threads)
        tBgB = gB_nkl[0, None, 0, k_tile, bidz]
        tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]

        # Use FP16 for A*B (reduces register pressure, better vectorization)
        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 and convert
        tArA.store(tAgA.load().to(cutlass.Float16))
        tBrB.store(tBgB.load().to(cutlass.Float16))
        tArSFA.store(tAgSFA.load().to(cutlass.Float32))
        tBrSFB.store(tBgSFB.load().to(cutlass.Float32))

        # Compute A*B first in FP16, then apply scales in FP32
        tABrAB = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
        tSFrSF = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
        tABrAB.store(tArA.load() * tBrB.load())
        tSFrSF.store(tArSFA.load() * tBrSFB.load())

        # Compute partial dot product for this K-tile with 8-way unrolling
        # NOTE: Loop bound now uses k_tile_size (128 or 256)
        # 8-way unrolling exposes more instruction-level parallelism (ILP)
        # Helps hide FMUL latency (~4 cycles) with independent operations
        for k_block in cutlass.range_constexpr(k_tile_size // 8):
            res += tABrAB[k_block*8+0] * tSFrSF[k_block*8+0]
            res += tABrAB[k_block*8+1] * tSFrSF[k_block*8+1]
            res += tABrAB[k_block*8+2] * tSFrSF[k_block*8+2]
            res += tABrAB[k_block*8+3] * tSFrSF[k_block*8+3]
            res += tABrAB[k_block*8+4] * tSFrSF[k_block*8+4]
            res += tABrAB[k_block*8+5] * tSFrSF[k_block*8+5]
            res += tABrAB[k_block*8+6] * tSFrSF[k_block*8+6]
            res += tABrAB[k_block*8+7] * tSFrSF[k_block*8+7]

    # NO ATOMICS! Each thread writes to unique (tidx, tidy) location
    shared_res[(tidx, tidy)] = res[0]
    cute.arch.sync_threads()

    # Explicit reduction over K dimension (only tidy==0 threads)
    # K-major layout ensures coalesced access (stride-1 in tidy)
    if tidy == 0:
        out = cute.zeros_like(tCgC, cutlass.Float32)

        # Reduce across all tidy values for this tidx
        for i in cutlass.range_constexpr(threads_per_k_param):
            out += shared_res[(tidx, i)]

        # Store final result (convert FP32 → FP16)
        tCgC.store(out.to(cutlass.Float16))

    return


@cute.jit
def cutedsl_path(
    a_ptr: cute.Pointer,
    b_ptr: cute.Pointer,
    sfa_ptr: cute.Pointer,
    sfb_ptr: cute.Pointer,
    c_ptr: cute.Pointer,
    problem_size: tuple,
    k_tile_size: cutlass.Constexpr,
    threads_per_k_param: cutlass.Constexpr,
):
    m, _, k, l = problem_size

    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)),
        ),
    )
    n_padded_128 = 128
    b_tensor = cute.make_tensor(
        b_ptr,
        cute.make_layout(
            (n_padded_128, cute.assume(k, 32), l),
            stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
        ),
    )
    c_tensor = cute.make_tensor(
        c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
    )
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
    sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
    sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)

    # Grid: 2D blocks for M and batch only (no K parallelism in grid)
    # K parallelism handled by 2D thread blocks (tidy dimension)
    grid = (
        cute.ceil_div(c_tensor.shape[0], threads_per_m),  # M blocks
        1,                                                 # Not used
        c_tensor.shape[2],                                 # Batch
    )

    # Pass k_tile_size and threads_per_k_param as parameters to kernel
    cutedsl_kernel(
        a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor, k_tile_size, threads_per_k_param
    ).launch(
        grid=grid,
        block=[threads_per_m, threads_per_k_param, 1],
        cluster=(1, 1, 1),
    )
    return


_compiled_cutedsl_cache = {}  # Changed to dict for multiple configs


def compile_cutedsl(k_tile_size, threads_per_k_val):
    """
    Compile kernel with specific K-tile size and threads_per_k.

    VECTORIZATION NOTE: We're going to work around CuteDSL limitations by using
    the fact that FP4 data is already packed in memory (2 elements per byte).
    We'll rely on the compiler's existing vectorization but verify it's working.
    """
    global _compiled_cutedsl_cache
    cache_key = (k_tile_size, threads_per_k_val)
    if cache_key in _compiled_cutedsl_cache:
        return _compiled_cutedsl_cache[cache_key]

    # Keep using FP4 dtype - CuteDSL should handle vectorization internally
    # but we'll verify with larger assumed_align to encourage wider loads
    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)
    c_ptr = make_ptr(c_dtype, 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)

    # Pass k_tile_size and threads_per_k_val as compile-time constants
    _compiled_cutedsl_cache[cache_key] = cute.compile(
        cutedsl_path, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0), k_tile_size, threads_per_k_val,
        options="--opt-level 3 --ptxas-options '--opt-level=3'"
    )

    return _compiled_cutedsl_cache[cache_key]


def custom_kernel(data: input_t) -> output_t:
    """
    Adaptive K-tiling and threads_per_k implementation.

    Key improvements:
    - Adaptive K-tile size: K=128 for small K, K=256 for large K
    - Adaptive threads_per_k: Reduces SMEM reduction overhead
      - k≤2048: threads_per_k=4 (17.4 µs on case 2)
      - k≤8192: threads_per_k=6 (50.0 µs on case 1)
      - k>8192: threads_per_k=8 (35.6 µs on case 0)
    """
    a, b, sfa_cpu, sfb_cpu, sfa_permuted, sfb_permuted, c = data

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

    # Determine optimal K-tile size and threads_per_k
    k_tile_size = get_k_tile_size(k)
    threads_per_k_val = get_threads_per_k(k)
    compiled_func = compile_cutedsl(k_tile_size, threads_per_k_val)

    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)
    c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    sfa_ptr = make_ptr(sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
    sfb_ptr = make_ptr(sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)

    compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, 1, k, l))

    return c

scrolls · 317 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