Skip to content
KernelIndex
Search⌘K

submission 92661

rt11 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-92661?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
280.1µs
#584 of 678
2025-11-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3a90ee38db2f86117048b1fa10df9f36cdd83211e4a6578e898690c191e48437
license declaredunknown
license concludedunknown
authorsrt11
imported2026-08-26

Techniques

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

fp4m.def("run_nvfp4_gemv", &run_nvfp4_gemv, "Warp-coalesced FP4 GEMV");
fp8const cutlass::float_e4m3_t* sfa_ptr,
shared-memory__shared__ __align__(16) __half2 sh_b[kPackedPerChunk];
vector-width = uint4__device__ __forceinline__ uint4 ld_uint4(const uint8_t* ptr) {

Kernel source

v1.py313 lines
# Warp-coalesced FP4 GEMV tuned for SM100 (Blackwell)
# v1: map one warp -> one row to fix uncoalesced global loads, keep shared B/SF
#     staging; vector-friendly strides and strict alignment for packed FP4.

import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

_module = None

# C++ glue
cpp_source = r"""
#include <cstdint>
#include <torch/extension.h>
#include "cutlass/float8.h"
#include "cutlass/half.h"

extern "C" void nvfp4_gemv_kernel_launch(
    const uint8_t* a_ptr,
    int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
    const uint8_t* b_ptr,
    int64_t stride_b_k, int64_t stride_b_l,
    const cutlass::float_e4m3_t* sfa_ptr,
    int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
    const cutlass::float_e4m3_t* sfb_ptr,
    int64_t stride_sfb_k, int64_t stride_sfb_l,
    cutlass::half_t* c_ptr,
    int64_t stride_c_m, int64_t stride_c_l,
    int32_t m, int32_t l, int32_t k_actual, int32_t k_blocks
);

torch::Tensor run_nvfp4_gemv(
    torch::Tensor a, torch::Tensor b,
    torch::Tensor sfa, torch::Tensor sfb,
    torch::Tensor c
) {
    constexpr int kSfVecSize = 16;

    TORCH_CHECK(a.is_cuda() && b.is_cuda() && sfa.is_cuda() && sfb.is_cuda() && c.is_cuda(),
                "All tensors must be on CUDA");
    TORCH_CHECK(a.dim() == 3 && b.dim() == 3 && sfa.dim() == 3 && sfb.dim() == 3 && c.dim() == 3,
                "All tensors must be 3-dimensional");

    const int64_t m = a.size(0);
    const int64_t packed_k = a.size(1);
    const int64_t l = a.size(2);

    TORCH_CHECK(m > 0 && packed_k > 0 && l > 0, "Dimensions must be positive");
    TORCH_CHECK(b.size(0) == 1 && b.size(1) == packed_k && b.size(2) == l,
                "B must be compacted to N=1 dimension");
    TORCH_CHECK(c.size(0) == m && c.size(1) == 1 && c.size(2) == l, "C shape mismatch");

    const int64_t k_actual = packed_k * 2;
    const int64_t k_blocks = (k_actual + kSfVecSize - 1) / kSfVecSize;

    TORCH_CHECK(sfa.size(0) == m && sfa.size(1) == k_blocks && sfa.size(2) == l, "SFA shape mismatch");
    TORCH_CHECK(sfb.size(0) == 1 && sfb.size(1) == k_blocks && sfb.size(2) == l,
                "SFB must be compacted to N=1 dimension");

    const int32_t m32 = static_cast<int32_t>(m);
    const int32_t l32 = static_cast<int32_t>(l);
    const int32_t k_actual32 = static_cast<int32_t>(k_actual);
    const int32_t k_blocks32 = static_cast<int32_t>(k_blocks);

    const int64_t a_bytes = a.element_size();
    const int64_t b_bytes = b.element_size();
    const int64_t sfa_bytes = sfa.element_size();
    const int64_t sfb_bytes = sfb.element_size();
    const int64_t c_bytes = c.element_size();

    nvfp4_gemv_kernel_launch(
        reinterpret_cast<const uint8_t*>(a.data_ptr()),
        a.stride(0) * a_bytes, a.stride(1) * a_bytes, a.stride(2) * a_bytes,
        reinterpret_cast<const uint8_t*>(b.data_ptr()),
        b.stride(1) * b_bytes, b.stride(2) * b_bytes,
        reinterpret_cast<const cutlass::float_e4m3_t*>(sfa.data_ptr()),
        sfa.stride(0) * sfa_bytes, sfa.stride(1) * sfa_bytes, sfa.stride(2) * sfa_bytes,
        reinterpret_cast<const cutlass::float_e4m3_t*>(sfb.data_ptr()),
        sfb.stride(1) * sfb_bytes, sfb.stride(2) * sfb_bytes,
        reinterpret_cast<cutlass::half_t*>(c.data_ptr()),
        c.stride(0) * c_bytes, c.stride(2) * c_bytes,
        m32, l32, k_actual32, k_blocks32
    );

    return c;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run_nvfp4_gemv", &run_nvfp4_gemv, "Warp-coalesced FP4 GEMV");
}
"""

# CUDA kernel
cuda_source = r"""
#include <algorithm>
#include <cstdint>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "cutlass/float8.h"
#include "cutlass/half.h"
#include "cutlass/numeric_conversion.h"

// Tunables
static constexpr int kLanesPerRow   = 32;  // full warp -> one row (coalesced)
static constexpr int kRowsPerWarp   = 1;
static constexpr int kWarpsPerBlock = 8;   // 8 rows per block
static constexpr int kThreads       = kWarpsPerBlock * 32;
static constexpr int kRowsPerBlock  = kRowsPerWarp * kWarpsPerBlock; // 8
static constexpr int kChunkElems    = 256; // elements per K tile
static constexpr int kPackedPerChunk= kChunkElems / 2;   // bytes
static constexpr int kSfVec         = 16;
static constexpr int kSfPerChunk    = kChunkElems / kSfVec; // 16

// Packed FP4x2 -> __half2 via CUTLASS converters (safe with ptxas)
__device__ __forceinline__ __half2 fp4x2_to_half2(const uint8_t v) {
    cutlass::float_e2m1_t lo = cutlass::float_e2m1_t::bitcast(v & 0xF);
    cutlass::float_e2m1_t hi = cutlass::float_e2m1_t::bitcast((v >> 4) & 0xF);
    cutlass::NumericConverter<cutlass::half_t, cutlass::float_e2m1_t> conv;
    cutlass::half_t h0 = conv(lo);
    cutlass::half_t h1 = conv(hi);
    return __halves2half2(reinterpret_cast<__half&>(h0), reinterpret_cast<__half&>(h1));
}

__device__ __forceinline__ __half fp8_to_half(const cutlass::float_e4m3_t v) {
    cutlass::NumericConverter<cutlass::half_t, cutlass::float_e4m3_t> conv;
    cutlass::half_t h = conv(v);
    return reinterpret_cast<__half&>(h);
}

// Vectorized byte load for A (16B) to guarantee 128B transactions across warp
__device__ __forceinline__ uint4 ld_uint4(const uint8_t* ptr) {
    return *reinterpret_cast<const uint4*>(__builtin_assume_aligned(ptr, 16));
}

extern "C" __global__ void __launch_bounds__(kThreads)
nvfp4_gemv_kernel_fast(
    const uint8_t* __restrict__ a_ptr,
    int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
    const uint8_t* __restrict__ b_ptr,
    int64_t stride_b_k, int64_t stride_b_l,
    const cutlass::float_e4m3_t* __restrict__ sfa_ptr,
    int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
    const cutlass::float_e4m3_t* __restrict__ sfb_ptr,
    int64_t stride_sfb_k, int64_t stride_sfb_l,
    cutlass::half_t* __restrict__ c_ptr,
    int64_t stride_c_m, int64_t stride_c_l,
    int32_t m, int32_t l, int32_t k_actual, int32_t /*k_blocks*/
) {
    // Batch selection
    const int batch = blockIdx.z;
    const uint8_t* a_base = a_ptr + batch * stride_a_l;
    const uint8_t* b_base = b_ptr + batch * stride_b_l;
    const uint8_t* sfa_base_bytes = reinterpret_cast<const uint8_t*>(sfa_ptr) + batch * stride_sfa_l;
    const uint8_t* sfb_base_bytes = reinterpret_cast<const uint8_t*>(sfb_ptr) + batch * stride_sfb_l;
    uint8_t* c_base_bytes = reinterpret_cast<uint8_t*>(c_ptr) + batch * stride_c_l;

    // Thread layout: one warp per row
    const int warp_id = threadIdx.x >> 5;   // 0..7
    const int lane_id = threadIdx.x & 31;   // 0..31

    const int block_row_start = blockIdx.x * kRowsPerBlock;
    const int global_row = block_row_start + warp_id;
    const bool active_row = (global_row < m);

    // Shared staging (align to 16B for vector ld/st)
    __shared__ __align__(16) __half2 sh_b[kPackedPerChunk];
    __shared__ __align__(16) cutlass::float_e4m3_t sh_sfb[kSfPerChunk];
    __shared__ __align__(16) cutlass::float_e4m3_t sh_sfa[kRowsPerBlock][kSfPerChunk];

    float acc = 0.f;

    for (int k_off = 0; k_off < k_actual; k_off += kChunkElems) {
        const int remaining = k_actual - k_off;
        const int chunk_elems = remaining > kChunkElems ? kChunkElems : remaining;
        const int packed_count = (chunk_elems + 1) / 2;
        const int sf_count = (chunk_elems + kSfVec - 1) / kSfVec;

        // Stage SFB for this chunk (first kThreads threads)
        for (int t = threadIdx.x; t < sf_count; t += kThreads) {
            const int sf_global = (k_off / kSfVec) + t;
            const auto* sf_ptr = reinterpret_cast<const cutlass::float_e4m3_t*>(
                sfb_base_bytes + sf_global * stride_sfb_k);
            sh_sfb[t] = *sf_ptr;
        }
        __syncthreads();

        // Stage B (packed FP4x2 -> half2, scaled)
        for (int t = threadIdx.x; t < packed_count; t += kThreads) {
            const int pb_global = (k_off / 2) + t;
            const uint8_t packed = *(b_base + pb_global * stride_b_k);
            __half2 hb = fp4x2_to_half2(packed);
            const int sf_idx = (t * 2) / kSfVec;
            const __half sf_h = fp8_to_half(sh_sfb[sf_idx]);
            const __half2 sf_h2 = __halves2half2(sf_h, sf_h);
            hb = __hmul2(hb, sf_h2);
            sh_b[t] = hb;
        }

        // Stage SFA for rows in this block (coalesced over M)
        const int rows_this_block = min(kRowsPerBlock, m - block_row_start);
        for (int t = threadIdx.x; t < rows_this_block * sf_count; t += kThreads) {
            const int local_row = t / sf_count;
            const int sf_local = t - local_row * sf_count;
            const int sf_global = (k_off / kSfVec) + sf_local;
            const auto* sfa_ptr = reinterpret_cast<const cutlass::float_e4m3_t*>(
                sfa_base_bytes + (block_row_start + local_row) * stride_sfa_m + sf_global * stride_sfa_k);
            sh_sfa[local_row][sf_local] = *sfa_ptr;
        }
        __syncthreads();

        // Compute for active rows: full-warp contiguous loads -> coalesced
        if (active_row) {
            const int sf_base = k_off / kSfVec;
            for (int pb = lane_id; pb < packed_count; pb += kLanesPerRow) {
                const int sf_idx = (pb * 2) / kSfVec;
                const __half sfa_h = fp8_to_half(sh_sfa[warp_id][sf_idx]);
                const __half2 sfa_h2 = __halves2half2(sfa_h, sfa_h);

                const int pb_global = (k_off / 2) + pb;
                const uint8_t packed_a = *(a_base + global_row * stride_a_m + pb_global * stride_a_k);
                __half2 ha = fp4x2_to_half2(packed_a);
                ha = __hmul2(ha, sfa_h2);

                const __half2 hb = sh_b[pb];
                float2 af = __half22float2(ha);
                float2 bf = __half22float2(hb);
                acc += af.x * bf.x + af.y * bf.y;
            }
        }
        __syncthreads();
    }

    // Reduce within warp
    for (int offset = 16; offset > 0; offset >>= 1) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }

    if (active_row && lane_id == 0) {
        auto* out_ptr = reinterpret_cast<cutlass::half_t*>(c_base_bytes + global_row * stride_c_m);
        *out_ptr = cutlass::half_t(acc);
    }
}

extern "C" void nvfp4_gemv_kernel_launch(
    const uint8_t* a_ptr,
    int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
    const uint8_t* b_ptr,
    int64_t stride_b_k, int64_t stride_b_l,
    const cutlass::float_e4m3_t* sfa_ptr,
    int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
    const cutlass::float_e4m3_t* sfb_ptr,
    int64_t stride_sfb_k, int64_t stride_sfb_l,
    cutlass::half_t* c_ptr,
    int64_t stride_c_m, int64_t stride_c_l,
    int32_t m, int32_t l, int32_t k_actual, int32_t k_blocks
) {
    dim3 grid((m + kRowsPerBlock - 1) / kRowsPerBlock, 1, l);
    dim3 block(kThreads);
    nvfp4_gemv_kernel_fast<<<grid, block>>>(
        a_ptr, stride_a_m, stride_a_k, stride_a_l,
        b_ptr, stride_b_k, stride_b_l,
        sfa_ptr, stride_sfa_m, stride_sfa_k, stride_sfa_l,
        sfb_ptr, stride_sfb_k, stride_sfb_l,
        c_ptr, stride_c_m, stride_c_l,
        m, l, k_actual, k_blocks
    );
}
"""


def _get_module():
    """Compile and cache the CUDA extension."""
    global _module
    if _module is None:
        repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))
        cutlass_include = os.path.join(repo_root, "cutlass", "include")
        _module = load_inline(
            name="nvfp4_gemv_v1_rowwarp",
            cpp_sources=[cpp_source],
            cuda_sources=[cuda_source],
            extra_include_paths=[cutlass_include],
            extra_cflags=["-std=c++20", "-O3"],
            extra_cuda_cflags=[
                "-std=c++20",
                "--use_fast_math",
                "-O3",
                "--expt-relaxed-constexpr",
                "-Xptxas=-v",
                "-arch=sm_100",
            ],
            verbose=False,
            with_cuda=True,
        )
    return _module


def custom_kernel(data: input_t) -> output_t:
    """
    Coalesced FP4 GEMV:
    - full warp owns a row (no cross-row mixing) to maximize global coalescing
    - shared decode of packed B and per-chunk scales
    - FP4x2 -> half2 conversion via cvt.rn.f16x2.e2m1x2
    """
    a, b_padded, sfa, sfb_padded, *_rest, c = data
    b = b_padded.narrow(0, 0, 1)
    sfb = sfb_padded.narrow(0, 0, 1)

    module = _get_module()
    module.run_nvfp4_gemv(a, b, sfa, sfb, c)
    return c
scrolls · 313 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