Skip to content
KernelIndex
Search⌘K

submission 108379

taka09203 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_gemv_cute__49MicroSec_opt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-108379?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
46.9µs
#259 of 678
2025-11-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:438b7d5d27d1604001cfa78f6515c8e7ba272df910b41cf20a190a06a1c5f7f1
license declaredunknown
license concludedunknown
authorstaka09203
imported2026-08-26

Techniques

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

fp4ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and B
split-kdef create_kernel_pipeline(num_split_k_const, tile_m=256):

Kernel source

nvfp4_gemv_cute__49MicroSec_opt.py381 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.blockscaled_layout as blockscaled_utils
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm, llvm
from cutlass import Float32

# Kernel configuration - base parameters
ab_dtype = cutlass.Float4E2M1FN  # FP4 data type for A and B
sf_dtype = cutlass.Float8E4M3FN  # FP8 data type for scale factors
c_dtype = cutlass.Float16        # FP16 output type
sf_vec_size = 16                 # Scale factor block size (16 elements share one scale)


# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


@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()
    )


def create_kernel_pipeline(num_split_k_const, tile_m=256):
    """
    Creates specialized GPU and JIT kernels for a specific Split-K factor.
    This avoids passing scalar arguments to the kernel, which CuTe DSL doesn't support well.
    """
    
    # Kernel configuration parameters optimized for NVIDIA B200 GPU
    # B200 features: 256KB L1 cache/SM, 50MB L2 cache, 6th-gen Tensor Cores with FP4 support
    # Optimizations: Increased thread count for better occupancy and latency hiding
    mma_tiler_mnk = (tile_m, 1, 128)  # Tile sizes for M, N, K dimensions (increased M tile for 256 threads)
    threads_per_cta = tile_m            # Number of threads per CUDA thread block (increased for better occupancy)

    # The CuTe reference implementation for NVFP4 block-scaled GEMV
    @cute.kernel
    def gpu_kernel(
        mA_mkl: cute.Tensor,
        mB_nkl: cute.Tensor,
        mSFA_mkl: cute.Tensor,
        mSFB_nkl: cute.Tensor,
        mC_mnl: cute.Tensor,
        mWorkspace_mnl: cute.Tensor,
    ):
        # Get CUDA block and thread indices
        bidx, bidy, bidz = cute.arch.block_idx()
        tidx, _, _ = cute.arch.thread_idx()

        # Extract the local tile for input matrix A (shape: [block_M, block_K, rest_M, rest_K, rest_L])
        gA_mkl = cute.local_tile(
            mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
        )
        # Extract the local tile for scale factor tensor for A (same shape as gA_mkl)
        # Here, block_M = (32, 4); block_K = (16, 4)
        gSFA_mkl = cute.local_tile(
            mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
        )
        # Extract the local tile for input matrix B (shape: [block_N, block_K, rest_N, rest_K, rest_L])
        gB_nkl = cute.local_tile(
            mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
        )
        # Extract the local tile for scale factor tensor for B (same shape as gB_nkl)
        gSFB_nkl = cute.local_tile(
            mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
        )
        # Extract the local tile for output matrix C (shape: [block_M, block_N, rest_M, rest_N, rest_L])
        gC_mnl = cute.local_tile(
            mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
        )
        
        # Extract local tile for Workspace (same shape as C)
        gWorkspace_mnl = cute.local_tile(
            mWorkspace_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
        )

        # Select output element corresponding to this thread and block indices
        tCgC = gC_mnl[tidx, None, bidx, 0, bidz] # bidy is used for split-k, so we use 0 for N-block
        tCgC = cute.make_tensor(tCgC.iterator, 1)
        
        # Workspace element
        tWgW = gWorkspace_mnl[tidx, None, bidx, 0, bidz]
        tWgW = cute.make_tensor(tWgW.iterator, 1)
        
        res = cute.zeros_like(tCgC, cutlass.Float32)

        # Get the number of k tiles (depth dimension) for the reduction loop
        total_k_tiles = gA_mkl.layout[3].shape
        
        # Split-K logic
        # bidy is the split index (0 to num_split_k - 1)
        # Use the captured constant num_split_k_const
        k_tiles_per_split = ceil_div(total_k_tiles, num_split_k_const)
        k_start = bidy * k_tiles_per_split
        k_end = min((bidy + 1) * k_tiles_per_split, total_k_tiles)

        for k_tile in range(k_start, k_end):
            tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
            tBgB = gB_nkl[0, None, 0, k_tile, bidz] # bidy is split-k, so use 0 for N-block
            tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
            tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]

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

            # Load NVFP4 or FP8 values from global memory
            a_val_nvfp4 = tAgA.load()
            b_val_nvfp4 = tBgB.load()
            sfa_val_fp8 = tAgSFA.load()
            sfb_val_fp8 = tBgSFB.load()

            # Convert loaded values to float32 for computation (FFMA)
            a_val = a_val_nvfp4.to(cutlass.Float32)
            b_val = b_val_nvfp4.to(cutlass.Float32)
            sfa_val = sfa_val_fp8.to(cutlass.Float32)
            sfb_val = sfb_val_fp8.to(cutlass.Float32)

            # Store the converted values to RMEM CuTe tensors
            tArA.store(a_val)
            tBrB.store(b_val)
            tArSFA.store(sfa_val)
            tBrSFB.store(sfb_val)

            # Iterate over SF vector tiles and compute the scale & matmul accumulation
            # Unrolled by factor of 8 for better ILP and reduced loop overhead on B200
            # Increased unrolling to maximize instruction-level parallelism
            k_dim = mma_tiler_mnk[2]  # 128
            for i in cutlass.range_constexpr(k_dim // 8):
                # Process 8 elements per iteration to maximize instruction-level parallelism
                base_idx = i * 8
                
                # Unrolled iterations (8-way unroll)
                scale0 = tArSFA[base_idx] * tBrSFB[base_idx]
                res += tArA[base_idx] * tBrB[base_idx] * scale0
                
                scale1 = tArSFA[base_idx + 1] * tBrSFB[base_idx + 1]
                res += tArA[base_idx + 1] * tBrB[base_idx + 1] * scale1
                
                scale2 = tArSFA[base_idx + 2] * tBrSFB[base_idx + 2]
                res += tArA[base_idx + 2] * tBrB[base_idx + 2] * scale2
                
                scale3 = tArSFA[base_idx + 3] * tBrSFB[base_idx + 3]
                res += tArA[base_idx + 3] * tBrB[base_idx + 3] * scale3
                
                scale4 = tArSFA[base_idx + 4] * tBrSFB[base_idx + 4]
                res += tArA[base_idx + 4] * tBrB[base_idx + 4] * scale4
                
                scale5 = tArSFA[base_idx + 5] * tBrSFB[base_idx + 5]
                res += tArA[base_idx + 5] * tBrB[base_idx + 5] * scale5
                
                scale6 = tArSFA[base_idx + 6] * tBrSFB[base_idx + 6]
                res += tArA[base_idx + 6] * tBrB[base_idx + 6] * scale6
                
                scale7 = tArSFA[base_idx + 7] * tBrSFB[base_idx + 7]
                res += tArA[base_idx + 7] * tBrB[base_idx + 7] * scale7

        # Store result
        if num_split_k_const > 1:
            # Atomic add to workspace (Float32)
            atomic_add_fp32(res[0], tWgW.iterator)
        else:
            # Standard store to Output (Float16)
            tCgC.store(res.to(cutlass.Float16))
        return

    @cute.jit
    def launch_kernel(
        a_ptr: cute.Pointer,
        b_ptr: cute.Pointer,
        sfa_ptr: cute.Pointer,
        sfb_ptr: cute.Pointer,
        c_ptr: cute.Pointer,
        workspace_ptr: cute.Pointer,
        problem_size: tuple,
    ):
        """
        Host-side JIT function to prepare tensors and launch GPU kernel.
        """
        m, _, k, l = problem_size
        # Create CuTe Tensor via pointer and 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)),
            ),
        )
        # We use n=128 to create the torch tensor to do fp4 computation via torch._scaled_mm
        # then copy torch tensor to cute tensor for cute customized kernel computation.
        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))
        )
        
        # Workspace tensor (Float32, same shape as C)
        workspace_tensor = cute.make_tensor(
            workspace_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
        )

        # Convert scale factor tensors to MMA layout
        # The layout matches Tensor Core requirements: (((32, 4), REST_M), ((SF_K, 4), REST_K), (1, REST_L))
        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)

        # Compute grid dimensions
        # Grid is (M_blocks, num_split_k, L)
        grid = (
            cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),  # block_M = 128
            num_split_k_const,
            c_tensor.shape[2],
        )

        # Launch the CUDA kernel
        gpu_kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor, workspace_tensor).launch(
            grid=grid,
            block=[threads_per_cta, 1, 1],
            cluster=(1, 1, 1),
        )
        return

    return launch_kernel


# Global cache for compiled kernels (keyed by (num_split_k, tile_m))
_compiled_kernel_cache = {}


def compile_kernel(num_split_k, tile_m=256):
    """
    Compile the kernel once and cache it.
    This should be called before any timing measurements.

    Args:
        num_split_k: Split-K factor
        tile_m: M tile size (128 or 256)

    Returns:
        The compiled kernel function
    """
    global _compiled_kernel_cache

    cache_key = (num_split_k, tile_m)
    if cache_key in _compiled_kernel_cache:
        return _compiled_kernel_cache[cache_key]

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    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)
    workspace_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)

    # Create the specialized kernel pipeline with specified tile_m
    launch_kernel = create_kernel_pipeline(num_split_k, tile_m)

    # Compile the kernel
    compiled_func = cute.compile(
        launch_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, workspace_ptr, (0, 0, 0, 0)
    )
    
    _compiled_kernel_cache[cache_key] = compiled_func
    return compiled_func


def custom_kernel(data: input_t) -> output_t:
    """
    Execute the block-scaled GEMV kernel.

    This is the main entry point called by the evaluation framework.
    It converts PyTorch tensors to CuTe tensors, launches the kernel,
    and returns the result.

    Args:
        data: Tuple of (a, b, sfa_cpu, sfb_cpu, sfa_permuted, sfb_permuted, c) PyTorch tensors
            a: [m, k, l] - Input matrix in float4e2m1fn
            b: [1, k, l] - Input vector in float4e2m1fn
            sfa_cpu: [m, k, l] - Scale factors in float8_e4m3fn (unused here)
            sfb_cpu: [1, k, l] - Scale factors in float8_e4m3fn (unused here)
            sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn
            sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
            c: [m, 1, l] - Output vector in float16

    Returns:
        Output tensor c with computed GEMV results
    """
    a, b, _, _, sfa_permuted, sfb_permuted, c = data

    # Get dimensions from MxKxL layout
    m, k, l = a.shape
    # Torch uses e2m1_x2 data type, thus k is halved
    k = k * 2
    # GEMV N dimension is always 1
    n = 1

    # Adaptive tile size selection
    # Use 256 threads for better occupancy if M is divisible by 256
    # Fall back to 128 threads for compatibility otherwise
    if m % 256 == 0:
        tile_m = 256
    else:
        tile_m = 128

    # Optimized Split-K heuristics for benchmarks
    # Based on analysis:
    # - Large K (16384): Split-K = 8 highly effective
    # - Medium K (7168): Split-K = 8 with large batch (l=8)
    # - Small K (2048): No Split-K (overhead > benefit)
    num_split_k = 1
    
    if k >= 12288:
        # Very large K: always use aggressive Split-K
        num_split_k = 8
    elif k >= 6144:
        # Medium-large K: use Split-K, especially with large batch
        if l >= 4:
            num_split_k = 8
        else:
            num_split_k = 4
    elif k >= 4096:
        # Medium K: moderate Split-K only for small batch
        if l <= 2:
            num_split_k = 4
        else:
            num_split_k = 1
    # else: k < 4096, no Split-K
    
    # Ensure kernel is compiled (will use cached version if available)
    compiled_func = compile_kernel(num_split_k, tile_m)
    
    # Allocate workspace if needed
    workspace = None
    workspace_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    
    if num_split_k > 1:
        # Workspace needs to be zero-initialized because we use atomic_add
        # NOTE: We need to ensure the layout matches what the kernel expects: stride=(1, 1, m)
        # torch.zeros((m, 1, l)) gives stride (l, l, 1) which is wrong for L > 1.
        # We create (l, 1, m) and transpose to get (m, 1, l) with stride (1, m, m) which matches (1, 1, m) effectively.
        workspace = torch.zeros((l, 1, m), dtype=torch.float32, device=a.device).transpose(0, 2)
        workspace_ptr = make_ptr(cutlass.Float32, workspace.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)

    # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
    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
    )

    # Execute the compiled kernel
    compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, workspace_ptr, (m, n, k, l))

    if num_split_k > 1:
        # Copy result from workspace to C
        c.copy_(workspace.to(dtype=torch.float16))

    return c
scrolls · 381 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