Skip to content
KernelIndex
Search⌘K

submission 91318

txacvalh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e5e0d831b394009d4f02f3761f3517476a209c6d8a948d74c673b0aa0b9d0a47
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__(BYTES_PER_THREAD_PER_K_TILE) uint32_t sB_u32[];
split-ktemplate<int L,int M,int K, int BLOCK_M, int THREADS_PER_BLOCK, int sB_size_in_bytes, bool SPLIT_K>
tile-k = 16constexpr int TILE_K = 16;
tile-m = 16constexpr int TILE_M = 16;
tile-n = 8constexpr int TILE_N = 8;

Kernel source

good_kernel.py385 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 fp8e4m3x2 = __nv_fp8x2_e4m3;
using fp8e4m3x4 = uint32_t;

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 SM_COUNT = 148;

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



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;
constexpr int BYTES_PER_THREAD_PER_K_TILE = 32;

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];
};
/* __device__ __forceinline__ fp8e4m3 pick_byte(uint32_t x, int byte_idx) {
    uint32_t tmp;
    // byte_idx in [0, 3], bit offset = byte_idx * 8, width = 8
    asm volatile (
        "bfe.u32 %0, %1, %2, 8;\n"
        : "=r"(tmp)
        : "r"(x), "r"(byte_idx * 8)
    );
    return static_cast<uint8_t>(tmp);
} */
__inline__ __device__
uint32_t dot_fp4x8(const uint32_t a, const uint32_t b, uint32_t acc) {

    asm volatile( \\
    "{\\n" \\
    ".reg .b8 byte0, byte1, byte2, byte3;\\n" \\
    ".reg .b8 byte4, byte5, byte6, byte7;\\n" \\
    ".reg .f16x2 cvt_0, cvt_1, cvt_2, cvt_3;\\n" \\
    ".reg .f16x2 cvt_4, cvt_5, cvt_6, cvt_7;\\n" \\
    "mov.b32 {byte0, byte1, byte2, byte3}, %1;\\n" \\
    "mov.b32 {byte4, byte5, byte6, byte7}, %2;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_0, byte0;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_1, byte1;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_2, byte2;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_3, byte3;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_4, byte4;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_5, byte5;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_6, byte6;\\n" \\
    "cvt.rn.f16x2.e2m1x2 cvt_7, byte7;\\n" \\
    "fma.rn.f16x2 %0, cvt_0, cvt_4, %0;\\n" \\
    "fma.rn.f16x2 %0, cvt_1, cvt_5, %0;\\n" \\
    "fma.rn.f16x2 %0, cvt_2, cvt_6, %0;\\n" \\
    "fma.rn.f16x2 %0, cvt_3, cvt_7, %0;\\n" \\
    "}\\n" : "+r"(acc) : "r"(a) , "r"(b));
        
        return acc;

}

// Warp-level reduction: sums 'val' across the active threads in the warp
__inline__ __device__
float warpReduceSum(float val) {

    // Do tree-reduction within the warp
    for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
        val += __shfl_down_sync(0xffffffff, val, offset);
    }
    return val;  // After this, lane 0 holds the sum of the warp
}

template<int L,int M,int K, int BLOCK_M, int THREADS_PER_BLOCK, int sB_size_in_bytes, bool SPLIT_K>
__global__ void __maxnreg__(32) custom_func_kernel(const uint8_t* const __restrict__ A, 
                           const uint8_t* const __restrict__ B, 
                           fp8e4m3 * __restrict__ scale_A, 
                           const uint32_t * const __restrict__ scale_B, 
                           half* const __restrict__ C ) {
    
    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__(BYTES_PER_THREAD_PER_K_TILE) uint32_t sB_u32[];
    auto sB_scale = reinterpret_cast<fp8e4m3x4 *>(
        reinterpret_cast<uint8_t *>(sB_u32) + sB_size_in_bytes);



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

    for (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {
        uint32_t rB_values[8];
        {
        void *ptr = (void *)(&B[b_l_base_off + i*BYTES_PER_THREAD_PER_K_TILE]);
        asm volatile( \\
        "{\\n" \\
        "ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\
        "}\\n" : "=r"(rB_values[0]), "=r"(rB_values[1]), "=r"(rB_values[2]), "=r"(rB_values[3]), "=r"(rB_values[4]), "=r"(rB_values[5]), "=r"(rB_values[6]), "=r"(rB_values[7]) : "l"(ptr));
        }
        int rest_idx = i / WARP_SIZE;

        #pragma unroll(8)
        for (int j = 0; j < 8; j++) {
        sB_u32[rest_idx * 256 + j * 32 + lane_id] = rB_values[j];
        }
    }

    int b_scale_bound = K / (16 * sizeof(uint32_t));
    auto curr_scale_B = &scale_B[(l_idx *  N_PADDED_128 * 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] = curr_scale_B[i];
    }


    __syncthreads();

    #pragma unroll
    for (;m_idx < M; m_idx += gridDim.x * BLOCK_M) {

    const int a_l_base_off = (l_idx * (M * K) + m_idx * (K)) / 2;
    constexpr int ATOM_SIZE = 128 * 4;
    const int TOTAL_NUM_ELEMENTS_PER_L = M * K / 16;
    const int rest_m = m_idx / 128;
    const int sub_m_idx1 = (m_idx % 128) / 32;
    const int sub_m_idx0 = m_idx % 32;

        
    const int scale_A_base = l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + m_idx * (K / 16);
    auto curr_scale_A = reinterpret_cast<fp8e4m3x4 * >(&scale_A[scale_A_base]);
    half2 acc{};
    
    #pragma unroll
    for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * BYTES_PER_THREAD_PER_K_TILE); tile_idx_k += WARP_SIZE) {
        constexpr int factor = (2 * BYTES_PER_THREAD_PER_K_TILE) / (16 * sizeof(fp8e4m3x4));
        uint32_t a_scales = curr_scale_A[tile_idx_k * factor];
        uint32_t b_scales = sB_scale[tile_idx_k];
        union {
            uint32_t u;
            fp8e4m3x2 f[2];
        } cvt_a, cvt_b;
        cvt_a.u = a_scales;
        cvt_b.u = b_scales;
        uint32_t rA_values[8];

        {
        void *ptr = (void *)(&A[a_l_base_off + tile_idx_k*BYTES_PER_THREAD_PER_K_TILE]);
        asm volatile( \\
        "{\\n" \\
        "ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\
        "}\\n" : "=r"(rA_values[0]), "=r"(rA_values[1]), "=r"(rA_values[2]), "=r"(rA_values[3]), "=r"(rA_values[4]), "=r"(rA_values[5]), "=r"(rA_values[6]), "=r"(rA_values[7]) : "l"(ptr));
        }



        const int base_off = (tile_idx_k / WARP_SIZE)*256 + lane_id;

        #pragma unroll(2)
        for (int i = 0; i < 8; i+=4) {
            uint32_t tile_acc2{};
            uint32_t tile_acc2_2{};

            uint32_t curr_sB[4];
            curr_sB[0] = sB_u32[base_off + 32 * i];
            curr_sB[1] = sB_u32[base_off + 32 * (i + 1)];
            curr_sB[2] = sB_u32[base_off + 32 * (i + 2)];
            curr_sB[3] = sB_u32[base_off + 32 * (i + 3)];

            tile_acc2 = dot_fp4x8(rA_values[i], curr_sB[0], tile_acc2);
            tile_acc2 = dot_fp4x8(rA_values[i+1], curr_sB[1],  tile_acc2);

            tile_acc2_2 = dot_fp4x8(rA_values[i+2], curr_sB[2],  tile_acc2_2);
            tile_acc2_2 = dot_fp4x8(rA_values[i+3], curr_sB[3],  tile_acc2_2);
            auto curr_scales = static_cast<half2>(cvt_a.f[i>>2]) * static_cast<half2>(cvt_b.f[i>>2]);
            
            half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);
            half2 tile_acc_2 = reinterpret_cast<__half2 &>(tile_acc2_2);

            acc = __hfma2(tile_acc, __half2(curr_scales.x, curr_scales.x), acc);
            acc = __hfma2(tile_acc_2, __half2(curr_scales.y, curr_scales.y), acc);

        }

    }
    half r = acc.x + acc.y;
    r = warpReduceSum(r);
    if (lane_id == 0) {
        C[l_idx * M + m_idx] = r;
    }
    

    }
}


#define LAUNCH_CUSTOM_CASE(LV, MV, KV)                                       \\
    if (L == (LV) && M == (MV) && K == (KV)) {                               \\
        constexpr int CONST_L = (LV);                                        \\
        constexpr int CONST_M = (MV);                                        \\
        constexpr int CONST_K = (KV);                                        \\
        constexpr int WARP_COUNT = CONST_L == 1 ? 8 : 8;                           \\
        constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;            \\
        constexpr int BLOCK_M = WARP_COUNT;                                  \\
        const int GRID_DIM_X = cuda::ceil_div(M, BLOCK_M);                   \\
        dim3 blocks(GRID_DIM_X, L, 1);                                       \\
        const int threads = THREADS_PER_BLOCK;                               \\
        constexpr int sB_block_size = 256*4;                             \\
        constexpr int sB_size_in_bytes = cuda::ceil_div(CONST_K / 2, sB_block_size) * sB_block_size;  \\
        constexpr int sB_scale_size_in_bytes = CONST_K / 16;                            \\
        constexpr int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes; \\
        custom_func_kernel<CONST_L, CONST_M, CONST_K,BLOCK_M, THREADS_PER_BLOCK,sB_size_in_bytes, false><<<blocks, threads,     \\
            shared_mem_size_in_bytes>>>(      \\
            reinterpret_cast<uint8_t *>(A.data_ptr()),                       \\
            reinterpret_cast<uint8_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()));                                      \\
    } else

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 blocks = (N + threads - 1) / threads;  




    LAUNCH_CUSTOM_CASE(1, 7168, 16384)
    LAUNCH_CUSTOM_CASE(8, 4096, 7168)
    LAUNCH_CUSTOM_CASE(4, 7168, 2048)
    {
        // final "else" block: unsupported shape
        TORCH_CHECK(false,
            "custom_func_kernel: unsupported shape L=", L,
            " M=", M, " K=", K);
    }
    

    /* 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+PTX;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, sfb, sfa_permuted, sfb_permuted, c_ref = data
    m, k, l = a_ref.shape
    k = k * 2

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

    if (l, m, k) in set([(1, 7168, 16384), (8, 4096, 7168), (4, 7168, 2048)]):
        args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa.permute(2, 0, 1), sfb.permute(2, 0, 1), 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)

        return c_ref
    else:
        tmp_c = torch.empty((l, 64, m), dtype=torch.float16, device="cuda")
        for l_idx in range(l):

            # (m, k) @ (n, k).T -> (m, n)
            torch._scaled_mm(
                a_ref[:, :, l_idx],
                b_ref[0:64, :, 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,
                out=tmp_c[l_idx,:,:].transpose(0,1),
            )
        return tmp_c[:, 0:1, :].permute(2, 1, 0)
scrolls · 385 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 85416.

⋯ 225 unchanged lines
half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);
half2 tile_acc_2 = reinterpret_cast<__half2 &>(tile_acc2_2);
- acc = __hfma2(__half2(static_cast<half>(tile_acc.x + tile_acc.y), static_cast<half>(tile_acc_2.x + tile_acc_2.y)), curr_scales, acc);
+ acc = __hfma2(tile_acc, __half2(curr_scales.x, curr_scales.x), acc);
+ acc = __hfma2(tile_acc_2, __half2(curr_scales.y, curr_scales.y), acc);
}

Best evidence level for this revision: reported

JSON