Skip to content
KernelIndex
Search⌘K

submission 79136

txacvalh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6342ba0c501edadfd1325bd1ebe6bed2b83588d04753f823cbcae561cbb3e33f
license declaredunknown
license concludedunknown
authorstxacvalh
imported2026-08-15

Techniques

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

fp4PyTorch reference implementation of NVFP4 block-scaled GEMV.
fp8using fp8e4m3 = __nv_fp8_e4m3;
shared-memoryextern __shared__ __align__(16) uint64_t sB_u64[];
tile-k = 16constexpr int TILE_K = 16;
tile-m = 16constexpr int TILE_M = 16;
tile-n = 8constexpr int TILE_N = 8;
vector-width = half2half2 rA_float4 = static_cast<half2>(rA_value.small_vec[i]);

Kernel source

kernel.py261 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t
import os

custom_func_cuda_source = """
// m16n8k64

using byte = uint8_t;
using int16 = int16_t;
using int32 = int32_t;
using int64 = int64_t;
using fp32 = float;
using fp16 = half;
using bf16 = nv_bfloat16;
using fp8e4m3 = __nv_fp8_e4m3;
using fp8e5m2 = __nv_fp8_e5m2;

using fp4e2m1 = __nv_fp4_e2m1;
using fp4e2m1x2 = __nv_fp4x2_e2m1;
using fp4e2m1x4 = __nv_fp4x4_e2m1;
using e8m0_t = uint8_t;

constexpr int WARP_SIZE = 32;
constexpr int WARP_COUNT = 16;
constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;
constexpr int SM_COUNT = 148;

constexpr int FP4X8_SIZE_IN_BYTES = 4;
constexpr int FP4X16_SIZE_IN_BYTES = 8;

constexpr int BLOCK_M = WARP_COUNT;
// constexpr int BLOCK_M = 128;


constexpr int TILE_M = 16;
constexpr int TILE_K = 16;
constexpr int TILE_N = 8;
constexpr int N_PADDED_128 = 128;
constexpr int N = 1;


template <typename T>
constexpr T DIVUP(const T &x, const T &y) {
    return (((x) + ((y)-1)) / (y));
}

__inline__ __device__
int scale_idx(int mn_idx, int k_idx, int l_idx, int MN, int K) {
    constexpr int ATOM_SIZE = 128 * 4;
    int TOTAL_NUM_ELEMENTS_PER_L = MN * K / 16;
    int rest_k = k_idx / 4;
    int sub_k_idx = k_idx % 4;
    int rest_mn = mn_idx / 128;
    int sub_mn_idx1 = (mn_idx % 128) / 32;
    int sub_mn_idx0 = mn_idx % 32;

    
    return l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + rest_mn * (128 * K / 16) + rest_k * ATOM_SIZE + sub_mn_idx0 * 16 + sub_mn_idx1 * 4 + sub_k_idx;
}

union fp4vec {
    uint64_t vec;
    fp4e2m1x2 small_vec[8];
};

// Warp-level reduction: sums 'val' across the active threads in the warp
__inline__ __device__
float warpReduceSum(float val) {
    // Mask of active lanes in this warp
    unsigned int mask = __activemask();

    // Do tree-reduction within the warp
    for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
        val += __shfl_down_sync(mask, val, offset);
    }
    return val;  // After this, lane 0 holds the sum of the warp
}
template <int GRID_DIM_X>
__global__ void custom_func_kernel(const uint64_t* const __restrict__ A, 
                           const uint64_t* const __restrict__ B, 
                           const fp8e4m3 * const __restrict__ scale_A, 
                           const uint32_t * const __restrict__ scale_B, 
                           half* const __restrict__ C, 
                           const int L, const int M, const int K, const int sB_size_in_bytes) {
    const int block_m_idx = blockIdx.x;
    const int l_idx = blockIdx.y;
    const int warp_id = threadIdx.x / WARP_SIZE;
    const int lane_id = threadIdx.x % WARP_SIZE;
    int m_idx = block_m_idx * BLOCK_M + warp_id;

    extern __shared__ __align__(16) uint64_t sB_u64[];
    fp8e4m3 * sB_scale = reinterpret_cast<fp8e4m3 *>(
        reinterpret_cast<uint8_t *>(sB_u64) + sB_size_in_bytes);



    // Load B into shared mem
    int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2 * sizeof(uint64_t));
    int b_bound = K / (2 * sizeof(uint64_t));

    #pragma unroll
    for (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {
        sB_u64[i] = B[b_l_base_off + i];
    }

    int b_scale_bound = K / (16 * sizeof(uint32_t));
    #pragma unroll
    for (int i = threadIdx.x; i < b_scale_bound; i += THREADS_PER_BLOCK) {
        reinterpret_cast<uint32_t *>(sB_scale)[i] = scale_B[scale_idx(0, i * sizeof(uint32_t), l_idx, N_PADDED_128, K) / sizeof(uint32_t)];
    }


    __syncthreads();

    for (;m_idx < M; m_idx += GRID_DIM_X * BLOCK_M) {

    const int a_l_base_off_u64 = (l_idx * (M * K) + m_idx * (K)) / (2 * sizeof(uint64_t));

    float acc{};

    #pragma unroll
    for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * sizeof(uint64_t)); tile_idx_k += WARP_SIZE) {

        fp4vec rA_value;
        fp4vec rB_value;

        rA_value.vec = A[a_l_base_off_u64 + tile_idx_k];
        rB_value.vec = sB_u64[tile_idx_k];
        half tile_acc{};


        #pragma unroll
        for (int i = 0; i < 8; i++) {
            half2 rA_float4 = static_cast<half2>(rA_value.small_vec[i]);
            half2 rB_float4 = static_cast<half2>(rB_value.small_vec[i]);
            tile_acc += rA_float4.x * rB_float4.x;
            tile_acc += rA_float4.y * rB_float4.y;
        }

        /* if (lane_id == 0) {
            //printf("hello");
            printf("L %d M %d K-tile %d A %d B %d acc %f %f %f %f\\n", l_idx, m_idx, tile_idx_k, rA_value.vec, rB_value.vec, tile_accs.x, tile_accs.y, tile_accs.z, tile_accs.w);
        } */

        const float a_scale = static_cast<float>(scale_A[scale_idx(m_idx, tile_idx_k * (2 * sizeof(uint64_t)) / 16, l_idx, M, K)]);
        // const float b_scale = static_cast<float>(sB_scale[tile_idx_k * (2 * sizeof(uint64_t)) / 16]);
        const float b_scale = static_cast<float>(sB_scale[tile_idx_k]);

        acc += static_cast<float>(tile_acc) * a_scale * b_scale;

    }
    acc = warpReduceSum(acc);
    if (lane_id == 0) {
        const int c_l_base_off_half = l_idx * (M) + m_idx;
        C[c_l_base_off_half] = static_cast<half>(acc);
    }

    }
}

void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C) {
    const int L = A.size(0);
    const int M = A.size(1);
    const int K = A.size(2) * 2;

    const int threads = THREADS_PER_BLOCK; 
    // const int blocks = (N + threads - 1) / threads;  

    const int sB_size_in_bytes = K / 2;
    const int sB_scale_size_in_bytes = K / 16;
    const int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes;


    constexpr int GRID_DIM_X = SM_COUNT;
    dim3 blocks(GRID_DIM_X, L, 1);    
    custom_func_kernel<GRID_DIM_X><<<blocks, threads, shared_mem_size_in_bytes>>>(
        reinterpret_cast<uint64_t *>(A.data_ptr()),
        reinterpret_cast<uint64_t *>(B.data_ptr()),
        reinterpret_cast<fp8e4m3 *>(scale_A.data_ptr()),
        reinterpret_cast<uint32_t *>(scale_B.data_ptr()),
        reinterpret_cast<half *>(C.data_ptr()),
        L, M, K, sB_size_in_bytes);

    

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

"""
# print(custom_func_cuda_source)
custom_func_cpp_source = """
#include <cuda.h>

#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_fp4.h>
#include <torch/extension.h>

void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C);
"""
os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0a;12.0a"

module = load_inline(
    name='custom_func',
    cpp_sources=custom_func_cpp_source,
    cuda_sources=custom_func_cuda_source,
    functions=['custom_func'],
    verbose=True,
    extra_cuda_cflags=['-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF_OPERATORS__ ', '-U__CUDA_NO_HALF2_OPERATORS__', '-O3', '-Xptxas', '-O3', '-use_fast_math'],
)


def custom_kernel(
    data: input_t,
) -> output_t:
    """
    Reference implementation of block-scale fp8 gemv
    Args:
        data: Tuple that expands to:
            a: torch.Tensor[float4e2m1fn] of shape [m, k, l],
            b: torch.Tensor[float4e2m1fn] of shape [1, k, l],
            sfa: torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l], used by reference implementation
            sfb: torch.Tensor[float8_e4m3fnuz] of shape [1, k // 16, l], used by reference implementation
            sfa_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_m, 4, rest_k, l],
            sfb_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_n, 4, rest_k, l],
            c: torch.Tensor[float16] of shape [m, 1, l]
    Returns:
        Tensor containing output in float16
        c: torch.Tensor[float16] of shape [m, 1, l]
    """
    """
    PyTorch reference implementation of NVFP4 block-scaled GEMV.
    """
    a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data
    m, k, l = a_ref.shape

    if l == 1:
        tmp_c = torch.empty((l, 64, m), dtype=torch.float16, device="cuda")

        torch._scaled_mm(
            a_ref[:, :, 0],
            b_ref[0:64, :, 0].transpose(0, 1),
            sfa_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
            sfb_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
            bias=None,
            out_dtype=torch.float16,
            out=tmp_c[0,:,:].transpose(0,1),
        )
        return tmp_c[:, 0:1, :].permute(2, 1, 0)
    else:
        args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())

        module.custom_func(*args)

        return c_ref
scrolls · 261 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 79075.

⋯ 237 unchanged lines
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data
+ m, k, l = a_ref.shape
- # Get dimensions from MxNxL layout
- # _, _, l = c_ref.shape
+ if l == 1:
+ tmp_c = torch.empty((l, 64, m), dtype=torch.float16, device="cuda")
- # # Call torch._scaled_mm to compute the GEMV result
- # for l_idx in range(l):
- # # Convert the scale factor tensor to blocked format
- # scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
- # scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])
- # # (m, k) @ (n, k).T -> (m, n)
- # res = torch._scaled_mm(
- # a_ref[:, :, l_idx],
- # b_ref[:, :, l_idx].transpose(0, 1),
- # scale_a.cuda(),
- # scale_b.cuda(),
- # bias=None,
- # out_dtype=torch.float16,
- # )
- # c_ref[:, 0, l_idx] = res[:, 0]
-
- # a = [l, m, k]
- # b = [l, 1, k]
- # [l, rest_n, rest_k, 32, 4, 4]
- # for e in [a_ref, b_ref, sfa_permuted, sfb_permuted, c_ref]:
- # print(e.stride(), e.shape)
- args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())
-
- # for a in args:
- # assert a.is_cuda, "All input tensors must be on GPU"
- # print(a.stride(), a.shape)
- module.custom_func(*args)
+ torch._scaled_mm(
+ a_ref[:, :, 0],
+ b_ref[0:64, :, 0].transpose(0, 1),
+ sfa_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
+ sfb_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
+ bias=None,
+ out_dtype=torch.float16,
+ out=tmp_c[0,:,:].transpose(0,1),
+ )
+ return tmp_c[:, 0:1, :].permute(2, 1, 0)
+ else:
+ args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())
- # for l_idx in range(l):
- # # (m, k) @ (n, k).T -> (m, n)
- # res = torch._scaled_mm(
- # a_ref[:, :, l_idx],
- # b_ref[:, :, l_idx].transpose(0, 1),
- # sfa_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),
- # sfb_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),
- # bias=None,
- # out_dtype=torch.float16,
- # )
- # c_ref[:, 0, l_idx] = res[:, 0]
- # module.custom_func
- return c_ref
No newline at end of file
+ module.custom_func(*args)
+
+ return c_ref
No newline at end of file
scrolls · 69 diff lines total

Best evidence level for this revision: reported

JSON