Skip to content
KernelIndex
Search⌘K

submission 79473

muxfd · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-79473?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
79.8µs
#357 of 678
2025-11-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c2a2a418a1c480c65bc943eb782a21e80ac9cac9ded17fa1b8029308a9e3047e
license declaredunknown
license concludedunknown
authorsmuxfd
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
fp8__device__ inline float fp8_to_float(__nv_fp8_e4m3 val) {
shared-memory__shared__ float warp_sums[WARPS_PER_ROW];

Kernel source

submission.py426 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 torch.utils.cpp_extension import load_inline

# Kernel configuration parameters
mma_tiler_mnk = (128, 1, 128)  # Tile sizes for M, N, K dimensions (larger K tile for big-K)
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)
threads_per_cta = 128  # Number of threads per CUDA thread block


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


# Lightweight CUDA kernel for large-K, single-batch use (avoids CuTe overhead)
nvfp4_gemv_cuda_source = r"""
#include <cuda_fp16.h>
#include <cuda_fp8.h>

__device__ __constant__ float FP4_E2M1_LUT[16] = {
    0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
    -0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
};

__device__ inline float unpack_fp4_low(uint8_t val) {
    return FP4_E2M1_LUT[val & 0x0F];
}

__device__ inline float unpack_fp4_high(uint8_t val) {
    return FP4_E2M1_LUT[(val >> 4) & 0x0F];
}

__device__ inline float fp8_to_float(__nv_fp8_e4m3 val) {
    return float(val);
}

template<int WARPS_PER_ROW>
__global__ void nvfp4_gemv_kernel(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const __nv_fp8_e4m3* __restrict__ scale_a,
    const __nv_fp8_e4m3* __restrict__ scale_b,
    __half* __restrict__ C,
    int M,
    int K,
    int L) {

    const int K_packed = K / 2;
    const int K_scale = K / 16;
    const int WARP_SIZE = 32;
    const int THREADS_PER_ROW = WARPS_PER_ROW * WARP_SIZE;

    int row = blockIdx.x;
    int batch = blockIdx.y;
    int tid_in_row = threadIdx.x;

    if (row >= M || batch >= L) return;

    float sum = 0.0f;

    const uint2* A_vec = (const uint2*)(A + (batch * M + row) * K_packed);
    const uint2* B_vec = (const uint2*)(B + batch * K_packed);

    int num_scale_blocks = K_scale;
    const __nv_fp8_e4m3* scale_a_row = scale_a + (batch * M + row) * K_scale;
    const __nv_fp8_e4m3* scale_b_row = scale_b + batch * K_scale;

    for (int scale_idx = tid_in_row; scale_idx < num_scale_blocks; scale_idx += THREADS_PER_ROW) {
        float sa = fp8_to_float(scale_a_row[scale_idx]);
        float sb = fp8_to_float(scale_b_row[scale_idx]);
        float scale = sa * sb;

        int vec_idx = scale_idx;
        uint2 a_data = A_vec[vec_idx];
        uint2 b_data = B_vec[vec_idx];

        uint8_t* a_bytes = (uint8_t*)&a_data;
        uint8_t* b_bytes = (uint8_t*)&b_data;

        float local_sum = 0.0f;
        #pragma unroll
        for (int i = 0; i < 8; i++) {
            float a0 = unpack_fp4_low(a_bytes[i]);
            float b0 = unpack_fp4_low(b_bytes[i]);
            float a1 = unpack_fp4_high(a_bytes[i]);
            float b1 = unpack_fp4_high(b_bytes[i]);
            local_sum += a0 * b0 + a1 * b1;
        }

        sum += scale * local_sum;
    }

    #pragma unroll
    for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
        sum += __shfl_down_sync(0xffffffff, sum, offset);
    }

    __shared__ float warp_sums[WARPS_PER_ROW];
    int warp_id = tid_in_row / WARP_SIZE;
    int lane_id = tid_in_row % WARP_SIZE;

    if (lane_id == 0) {
        warp_sums[warp_id] = sum;
    }
    __syncthreads();

    if (warp_id == 0 && lane_id < WARPS_PER_ROW) {
        sum = warp_sums[lane_id];
        #pragma unroll
        for (int offset = WARPS_PER_ROW / 2; offset > 0; offset /= 2) {
            sum += __shfl_down_sync(0xffffffff, sum, offset);
        }
        if (lane_id == 0) {
            C[batch * M + row] = __float2half(sum);
        }
    }
}

void nvfp4_gemv_cuda(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor scale_a,
    torch::Tensor scale_b,
    torch::Tensor C,
    int M,
    int K,
    int L) {

    int warps_per_row;
    if (K >= 8192) {
        warps_per_row = 8;
    } else if (K >= 4096) {
        warps_per_row = 4;
    } else if (K >= 1024) {
        warps_per_row = 2;
    } else {
        warps_per_row = 1;
    }

    int threads_per_block = warps_per_row * 32;
    dim3 threads(threads_per_block);
    dim3 blocks(M, L);

    if (warps_per_row == 8) {
        nvfp4_gemv_kernel<8><<<blocks, threads>>>(
            reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
            reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
            M, K, L);
    } else if (warps_per_row == 4) {
        nvfp4_gemv_kernel<4><<<blocks, threads>>>(
            reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
            reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
            M, K, L);
    } else if (warps_per_row == 2) {
        nvfp4_gemv_kernel<2><<<blocks, threads>>>(
            reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
            reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
            M, K, L);
    } else {
        nvfp4_gemv_kernel<1><<<blocks, threads>>>(
            reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
            reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
            reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
            M, K, L);
    }

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
}
"""

nvfp4_gemv_cpp_source = r"""
void nvfp4_gemv_cuda(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor scale_a,
    torch::Tensor scale_b,
    torch::Tensor C,
    int M,
    int K,
    int L);
"""

nvfp4_cuda_module = load_inline(
    name='nvfp4_gemv_hybrid_cuda',
    cpp_sources=nvfp4_gemv_cpp_source,
    cuda_sources=nvfp4_gemv_cuda_source,
    functions=['nvfp4_gemv_cuda'],
    verbose=False,
    extra_cuda_cflags=['-O3', '--use_fast_math'],
)


# The CuTe implementation for NVFP4 block-scaled GEMV
@cute.kernel
def kernel(
    mA_mkl: cute.Tensor,
    mB_nkl: cute.Tensor,
    mSFA_mkl: cute.Tensor,
    mSFB_nkl: cute.Tensor,
    mC_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)
    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)
    )

    # Select output element corresponding to this thread and block indices
    tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
    tCgC = cute.make_tensor(tCgC.iterator, 1)
    res = cute.zeros_like(tCgC, cutlass.Float32)

    # Get the number of k tiles (depth dimension) for the reduction loop
    k_tile_cnt = gA_mkl.layout[3].shape
    for k_tile in range(k_tile_cnt):
        tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
        tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
        tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
        tBgSFB = gSFB_nkl[0, None, bidy, 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
        for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
            res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]

    # Store the final float16 result back to global memory
    tCgC.store(res.to(cutlass.Float16))
    return


@cute.jit
def my_kernel(
    a_ptr: cute.Pointer,
    b_ptr: cute.Pointer,
    sfa_ptr: cute.Pointer,
    sfb_ptr: cute.Pointer,
    c_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 customize kernel computation
    # therefore we need to ensure b_tensor has the right stride with this 128 padded size on n.
    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))
    )
    # 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)

    # Grid is (M_blocks, 1, L)
    grid = (
        cute.ceil_div(c_tensor.shape[0], 128),
        1,
        c_tensor.shape[2],
    )

    # Launch the CUDA kernel
    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


# Global cache for compiled kernel
_compiled_kernel_cache = None


# Compile the kernel once and cache it.
def compile_kernel():
    global _compiled_kernel_cache

    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache

    # 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)

    # Compile the kernel
    _compiled_kernel_cache = cute.compile(
        my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)
    )

    return _compiled_kernel_cache


def custom_kernel(data: input_t) -> output_t:
    """
    Hybrid: use CUDA path for large-K single batch; CuTe path otherwise.
    """
    a, b, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c = data

    # Get dimensions from MxKxL layout
    m, k_packed, l = a.shape
    k = k_packed * 2  # torch uses e2m1_x2 so K is packed by 2

    # Large-K single batch: use CUDA kernel (faster for K>=8192, l==1)
    if l == 1 and k >= 8192:
        a_uint8 = a.contiguous().view(torch.uint8)
        b_uint8 = b[0:1, :, :].contiguous().view(torch.uint8)
        sfa = sfa_ref.contiguous()
        sfb = sfb_ref[0:1, :, :].contiguous()
        c_slice = c[:, 0, 0].contiguous()

        nvfp4_cuda_module.nvfp4_gemv_cuda(
            a_uint8,
            b_uint8,
            sfa,
            sfb,
            c_slice,
            m,
            k,
            l,
        )
        c[:, 0, 0] = c_slice
        return c

    # Default: CuTe kernel
    compiled_func = compile_kernel()

    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 · 426 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