Skip to content
KernelIndex
Search⌘K

submission 110745

rcmalli · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-110745?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
26.0µs
#111 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f6cf2a1631537d4efc6beb45e33fbcdc323d166a6837193f7cd64e70aeff6c78
license declaredunknown
license concludedunknown
authorsrcmalli
imported2026-08-15

Techniques

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

fp4NVFP4 GEMV Kernel - v9: Adaptive Dispatch (SMEM B vs Direct)
fp8const __nv_fp8_e4m3* __restrict__ SFA,
shared-memorygemv_kernel_smem_b(
vector-width = uint4uint4 a_vec, uint4 b_vec,

Kernel source

submission_v9.py356 lines
"""
NVFP4 GEMV Kernel - v9: Adaptive Dispatch (SMEM B vs Direct)

Combines best of v8 and direct approach:
- SMEM B caching for large K + single batch (Case 1): 26.1µs
- Direct loads for batched cases (Cases 2 & 3): ~49µs, ~20µs

Expected geomean: ~30µs (vs v3's 32.0µs = 6% improvement)
"""

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

cuda_src = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <torch/extension.h>

using at::Tensor;

// Process 32 FP4 elements with 2 scale factors using PTX
__device__ __forceinline__
float process_32_fp4_with_scale_ptx(
    uint4 a_vec, uint4 b_vec,
    uint16_t sfa_packed, uint16_t sfb_packed)
{
    float result;
    asm volatile(
        "{\n"
        ".reg .b8 a0, a1, a2, a3, a4, a5, a6, a7;\n"
        ".reg .b8 a8, a9, a10, a11, a12, a13, a14, a15;\n"
        ".reg .b8 b0, b1, b2, b3, b4, b5, b6, b7;\n"
        ".reg .b8 b8, b9, b10, b11, b12, b13, b14, b15;\n"
        ".reg .f16x2 cvtA0, cvtA1, cvtA2, cvtA3, cvtA4, cvtA5, cvtA6, cvtA7;\n"
        ".reg .f16x2 cvtA8, cvtA9, cvtA10, cvtA11, cvtA12, cvtA13, cvtA14, cvtA15;\n"
        ".reg .f16x2 cvtB0, cvtB1, cvtB2, cvtB3, cvtB4, cvtB5, cvtB6, cvtB7;\n"
        ".reg .f16x2 cvtB8, cvtB9, cvtB10, cvtB11, cvtB12, cvtB13, cvtB14, cvtB15;\n"
        ".reg .f16x2 acc0_0, acc0_1, acc0_2, acc0_3;\n"
        ".reg .f16x2 acc1_0, acc1_1, acc1_2, acc1_3;\n"
        ".reg .f16x2 sfA_f16x2, sfB_f16x2, sf_f16x2;\n"
        ".reg .f16 sf0, sf1, lane0, lane1;\n"
        ".reg .f32 result_f32, tmp_f32;\n"
        "mov.b32 acc0_0, 0;\n" "mov.b32 acc0_1, 0;\n"
        "mov.b32 acc0_2, 0;\n" "mov.b32 acc0_3, 0;\n"
        "mov.b32 acc1_0, 0;\n" "mov.b32 acc1_1, 0;\n"
        "mov.b32 acc1_2, 0;\n" "mov.b32 acc1_3, 0;\n"
        "mov.b32 {a0, a1, a2, a3}, %1;\n"
        "mov.b32 {a4, a5, a6, a7}, %2;\n"
        "mov.b32 {a8, a9, a10, a11}, %3;\n"
        "mov.b32 {a12, a13, a14, a15}, %4;\n"
        "mov.b32 {b0, b1, b2, b3}, %5;\n"
        "mov.b32 {b4, b5, b6, b7}, %6;\n"
        "mov.b32 {b8, b9, b10, b11}, %7;\n"
        "mov.b32 {b12, b13, b14, b15}, %8;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA0, a0;\n" "cvt.rn.f16x2.e2m1x2 cvtA1, a1;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA2, a2;\n" "cvt.rn.f16x2.e2m1x2 cvtA3, a3;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA4, a4;\n" "cvt.rn.f16x2.e2m1x2 cvtA5, a5;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA6, a6;\n" "cvt.rn.f16x2.e2m1x2 cvtA7, a7;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA8, a8;\n" "cvt.rn.f16x2.e2m1x2 cvtA9, a9;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA10, a10;\n" "cvt.rn.f16x2.e2m1x2 cvtA11, a11;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA12, a12;\n" "cvt.rn.f16x2.e2m1x2 cvtA13, a13;\n"
        "cvt.rn.f16x2.e2m1x2 cvtA14, a14;\n" "cvt.rn.f16x2.e2m1x2 cvtA15, a15;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB0, b0;\n" "cvt.rn.f16x2.e2m1x2 cvtB1, b1;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB2, b2;\n" "cvt.rn.f16x2.e2m1x2 cvtB3, b3;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB4, b4;\n" "cvt.rn.f16x2.e2m1x2 cvtB5, b5;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB6, b6;\n" "cvt.rn.f16x2.e2m1x2 cvtB7, b7;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB8, b8;\n" "cvt.rn.f16x2.e2m1x2 cvtB9, b9;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB10, b10;\n" "cvt.rn.f16x2.e2m1x2 cvtB11, b11;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB12, b12;\n" "cvt.rn.f16x2.e2m1x2 cvtB13, b13;\n"
        "cvt.rn.f16x2.e2m1x2 cvtB14, b14;\n" "cvt.rn.f16x2.e2m1x2 cvtB15, b15;\n"
        "fma.rn.f16x2 acc0_0, cvtA0, cvtB0, acc0_0;\n"
        "fma.rn.f16x2 acc0_0, cvtA1, cvtB1, acc0_0;\n"
        "fma.rn.f16x2 acc0_1, cvtA2, cvtB2, acc0_1;\n"
        "fma.rn.f16x2 acc0_1, cvtA3, cvtB3, acc0_1;\n"
        "fma.rn.f16x2 acc0_2, cvtA4, cvtB4, acc0_2;\n"
        "fma.rn.f16x2 acc0_2, cvtA5, cvtB5, acc0_2;\n"
        "fma.rn.f16x2 acc0_3, cvtA6, cvtB6, acc0_3;\n"
        "fma.rn.f16x2 acc0_3, cvtA7, cvtB7, acc0_3;\n"
        "fma.rn.f16x2 acc1_0, cvtA8, cvtB8, acc1_0;\n"
        "fma.rn.f16x2 acc1_0, cvtA9, cvtB9, acc1_0;\n"
        "fma.rn.f16x2 acc1_1, cvtA10, cvtB10, acc1_1;\n"
        "fma.rn.f16x2 acc1_1, cvtA11, cvtB11, acc1_1;\n"
        "fma.rn.f16x2 acc1_2, cvtA12, cvtB12, acc1_2;\n"
        "fma.rn.f16x2 acc1_2, cvtA13, cvtB13, acc1_2;\n"
        "fma.rn.f16x2 acc1_3, cvtA14, cvtB14, acc1_3;\n"
        "fma.rn.f16x2 acc1_3, cvtA15, cvtB15, acc1_3;\n"
        "add.rn.f16x2 acc0_0, acc0_0, acc0_1;\n"
        "add.rn.f16x2 acc0_2, acc0_2, acc0_3;\n"
        "add.rn.f16x2 acc0_0, acc0_0, acc0_2;\n"
        "add.rn.f16x2 acc1_0, acc1_0, acc1_1;\n"
        "add.rn.f16x2 acc1_2, acc1_2, acc1_3;\n"
        "add.rn.f16x2 acc1_0, acc1_0, acc1_2;\n"
        "cvt.rn.f16x2.e4m3x2 sfA_f16x2, %9;\n"
        "cvt.rn.f16x2.e4m3x2 sfB_f16x2, %10;\n"
        "mul.rn.f16x2 sf_f16x2, sfA_f16x2, sfB_f16x2;\n"
        "mov.b32 {sf0, sf1}, sf_f16x2;\n"
        "mov.b32 {lane0, lane1}, acc0_0;\n"
        "mul.rn.f16 lane0, lane0, sf0;\n"
        "mul.rn.f16 lane1, lane1, sf0;\n"
        "add.rn.f16 lane0, lane0, lane1;\n"
        "cvt.f32.f16 result_f32, lane0;\n"
        "mov.b32 {lane0, lane1}, acc1_0;\n"
        "mul.rn.f16 lane0, lane0, sf1;\n"
        "mul.rn.f16 lane1, lane1, sf1;\n"
        "add.rn.f16 lane0, lane0, lane1;\n"
        "cvt.f32.f16 tmp_f32, lane0;\n"
        "add.f32 result_f32, result_f32, tmp_f32;\n"
        "mov.f32 %0, result_f32;\n"
        "}\n"
        : "=f"(result)
        : "r"(a_vec.x), "r"(a_vec.y), "r"(a_vec.z), "r"(a_vec.w),
          "r"(b_vec.x), "r"(b_vec.y), "r"(b_vec.z), "r"(b_vec.w),
          "h"(sfa_packed), "h"(sfb_packed)
    );
    return result;
}

//=============================================================================
// Kernel 1: SMEM B caching (for large K + single batch)
//=============================================================================
template<int WARPS_PER_BLOCK>
__global__ void __launch_bounds__(32 * WARPS_PER_BLOCK, 4)
gemv_kernel_smem_b(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const __nv_fp8_e4m3* __restrict__ SFA,
    const __nv_fp8_e4m3* __restrict__ SFB,
    __half* __restrict__ C,
    int M, int K, int L,
    int64_t sAm, int64_t sAk, int64_t sAl,
    int64_t sBm, int64_t sBk, int64_t sBl,
    int64_t sSFAm, int64_t sSFAk, int64_t sSFAl,
    int64_t sSFBm, int64_t sSFBk, int64_t sSFBl,
    int64_t sCm, int64_t sCn, int64_t sCl
) {
    extern __shared__ uint8_t smem[];

    const int K_packed = K / 2;
    const int num_sf = K / 16;

    uint8_t* smem_b = smem;
    uint8_t* smem_sfb = smem + K_packed;

    const int lane = threadIdx.x;
    const int warp_in_blk = threadIdx.y;
    const int tid = warp_in_blk * 32 + lane;
    const int block_threads = WARPS_PER_BLOCK * 32;

    const int row = blockIdx.x * WARPS_PER_BLOCK + warp_in_blk;
    const int batch = blockIdx.y;

    const uint8_t* B_row = B + (int64_t)batch * sBl;
    const __nv_fp8_e4m3* SFB_row = SFB + (int64_t)batch * sSFBl;

    // Cooperative load B into SMEM
    for (int i = tid * 16; i < K_packed; i += block_threads * 16) {
        if (i + 16 <= K_packed) {
            *reinterpret_cast<uint4*>(smem_b + i) =
                *reinterpret_cast<const uint4*>(B_row + i);
        } else {
            for (int j = i; j < min(i + 16, K_packed); j++) {
                smem_b[j] = B_row[j];
            }
        }
    }

    // Cooperative load SFB into SMEM
    for (int i = tid * 16; i < num_sf; i += block_threads * 16) {
        if (i + 16 <= num_sf) {
            *reinterpret_cast<uint4*>(smem_sfb + i) =
                *reinterpret_cast<const uint4*>(
                    reinterpret_cast<const uint8_t*>(SFB_row) + i * sSFBk);
        } else {
            for (int j = i; j < min(i + 16, num_sf); j++) {
                smem_sfb[j] = reinterpret_cast<const uint8_t*>(SFB_row)[j * sSFBk];
            }
        }
    }

    __syncthreads();

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

    const int num_sf_blocks = K / 16;
    const int sf_pairs = num_sf_blocks / 2;
    float acc = 0.0f;

    const uint8_t* A_row = A + (int64_t)row * sAm + (int64_t)batch * sAl;
    const __nv_fp8_e4m3* SFA_row = SFA + (int64_t)row * sSFAm + (int64_t)batch * sSFAl;

    for (int pair_idx = lane; pair_idx < sf_pairs; pair_idx += 32) {
        int k_packed_base = pair_idx * 16;
        int sf_idx_base = pair_idx * 2;

        uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_packed_base);
        uint4 b_vec = *reinterpret_cast<const uint4*>(smem_b + k_packed_base);

        uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(
            reinterpret_cast<const uint8_t*>(SFA_row) + sf_idx_base * sSFAk);
        uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(smem_sfb + sf_idx_base);

        acc += process_32_fp4_with_scale_ptx(a_vec, b_vec, sfa_packed, sfb_packed);
    }

    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        acc += __shfl_down_sync(0xffffffffu, acc, offset);
    }

    if (lane == 0) {
        int64_t c_idx = (int64_t)row * sCm + (int64_t)0 * sCn + (int64_t)batch * sCl;
        C[c_idx] = __float2half(acc);
    }
}

//=============================================================================
// Kernel 2: Direct loads (for batched / small K cases)
//=============================================================================
template<int WARPS_PER_BLOCK>
__global__ void gemv_kernel_direct(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const __nv_fp8_e4m3* __restrict__ SFA,
    const __nv_fp8_e4m3* __restrict__ SFB,
    __half* __restrict__ C,
    int M, int K, int L,
    int64_t sAm, int64_t sAk, int64_t sAl,
    int64_t sBm, int64_t sBk, int64_t sBl,
    int64_t sSFAm, int64_t sSFAk, int64_t sSFAl,
    int64_t sSFBm, int64_t sSFBk, int64_t sSFBl,
    int64_t sCm, int64_t sCn, int64_t sCl
) {
    const int lane = threadIdx.x;
    const int warp_in_blk = threadIdx.y;
    const int row = blockIdx.x * WARPS_PER_BLOCK + warp_in_blk;
    const int batch = blockIdx.y;

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

    const int num_sf_blocks = K / 16;
    const int sf_pairs = num_sf_blocks / 2;
    float acc = 0.0f;

    const uint8_t* A_row = A + (int64_t)row * sAm + (int64_t)batch * sAl;
    const uint8_t* B_row = B + (int64_t)batch * sBl;
    const __nv_fp8_e4m3* SFA_row = SFA + (int64_t)row * sSFAm + (int64_t)batch * sSFAl;
    const __nv_fp8_e4m3* SFB_row = SFB + (int64_t)batch * sSFBl;

    for (int pair_idx = lane; pair_idx < sf_pairs; pair_idx += 32) {
        int k_packed_base = pair_idx * 16;
        int sf_idx_base = pair_idx * 2;

        uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_packed_base);
        uint4 b_vec = *reinterpret_cast<const uint4*>(B_row + k_packed_base);

        uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(
            reinterpret_cast<const uint8_t*>(SFA_row) + sf_idx_base * sSFAk);
        uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(
            reinterpret_cast<const uint8_t*>(SFB_row) + sf_idx_base * sSFBk);

        acc += process_32_fp4_with_scale_ptx(a_vec, b_vec, sfa_packed, sfb_packed);
    }

    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        acc += __shfl_down_sync(0xffffffffu, acc, offset);
    }

    if (lane == 0) {
        int64_t c_idx = (int64_t)row * sCm + (int64_t)0 * sCn + (int64_t)batch * sCl;
        C[c_idx] = __float2half(acc);
    }
}

//=============================================================================
// Adaptive dispatch
//=============================================================================
void run_gemv(Tensor a, Tensor b, Tensor sfa, Tensor sfb, Tensor c) {
    const int64_t M = a.size(0);
    const int64_t K_packed = a.size(1);
    const int64_t K = K_packed * 2;
    const int64_t L = a.size(2);

    const uint8_t* A_ptr = reinterpret_cast<const uint8_t*>(a.data_ptr());
    const uint8_t* B_ptr = reinterpret_cast<const uint8_t*>(b.data_ptr());
    const __nv_fp8_e4m3* SFA = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());
    const __nv_fp8_e4m3* SFB = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());
    __half* C = reinterpret_cast<__half*>(c.data_ptr());

    auto a_str = a.strides();
    auto b_str = b.strides();
    auto sfa_str = sfa.strides();
    auto sfb_str = sfb.strides();
    auto c_str = c.strides();

    constexpr int WARPS_PER_BLOCK = 4;
    dim3 block(32, WARPS_PER_BLOCK, 1);
    dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L, 1);

    // Adaptive dispatch: SMEM B for large K + single batch, direct otherwise
    // SMEM B benefits when B is reused across many rows in same block
    // For batched (L > 1), each batch needs separate B, reducing SMEM benefit
    bool use_smem_b = (K >= 12000) && (L == 1);

    if (use_smem_b) {
        size_t smem_size = K_packed + (K / 16);
        gemv_kernel_smem_b<WARPS_PER_BLOCK>
            <<<grid, block, smem_size>>>(
                A_ptr, B_ptr, SFA, SFB, C,
                (int)M, (int)K, (int)L,
                a_str[0], a_str[1], a_str[2],
                b_str[0], b_str[1], b_str[2],
                sfa_str[0], sfa_str[1], sfa_str[2],
                sfb_str[0], sfb_str[1], sfb_str[2],
                c_str[0], c_str[1], c_str[2]);
    } else {
        gemv_kernel_direct<WARPS_PER_BLOCK>
            <<<grid, block>>>(
                A_ptr, B_ptr, SFA, SFB, C,
                (int)M, (int)K, (int)L,
                a_str[0], a_str[1], a_str[2],
                b_str[0], b_str[1], b_str[2],
                sfa_str[0], sfa_str[1], sfa_str[2],
                sfb_str[0], sfb_str[1], sfb_str[2],
                c_str[0], c_str[1], c_str[2]);
    }
}
"""

cpp_src = r"""
#include <torch/extension.h>
void run_gemv(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""

gemv_lib = load_inline(
    name="fp4_gemv_v9_adaptive",
    cpp_sources=cpp_src,
    cuda_sources=cuda_src,
    functions=["run_gemv"],
    extra_cuda_cflags=[
        "-O3",
        "-use_fast_math",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--expt-relaxed-constexpr",
    ],
)


def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
    gemv_lib.run_gemv(a, b, sfa, sfb, c)
    return c
scrolls · 356 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