Skip to content
KernelIndex
Search⌘K

submission 107290

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

kmajor.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-107290?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
30.9µs
#162 of 678
2025-11-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:50f4fcd9b215c65359e926ac01c26aa547b44156c95995cf8d075e5751fc13f3
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15

Techniques

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

vector-width = half2half2 acc_h2 = __float2half2_rn(0.0f);

Kernel source

kmajor.py154 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

cuda_source = """
#include <cuda_fp16.h>
#include <cuda_fp8.h>

__device__ __forceinline__ float warp_reduce_sum(float val) {
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1)
        val += __shfl_xor_sync(0xffffffff, val, offset);
    return val;
}

__device__ __forceinline__ void cvt_f4x8_to_f16x8(uint32_t src, uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3) {
    asm volatile(
        "{ .reg .b8 b0, b1, b2, b3; "
        "mov.b32 {b0, b1, b2, b3}, %4; "
        "cvt.rn.f16x2.e2m1x2 %0, b0; "
        "cvt.rn.f16x2.e2m1x2 %1, b1; "
        "cvt.rn.f16x2.e2m1x2 %2, b2; "
        "cvt.rn.f16x2.e2m1x2 %3, b3; }"
        : "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)
        : "r"(src)
    );
}

__device__ __forceinline__ uint32_t cvt_f8x2_to_f16x2(uint16_t src) {
    uint32_t dst;
    asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(dst) : "h"(src));
    return dst;
}

__global__ void gemv_v8(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    half* __restrict__ C,
    int M, int K, int L,
    int64_t a_s0, int64_t a_s2, int64_t b_s2,
    int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
    int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
    int64_t c_s0, int64_t c_s2
) {
    constexpr int WARPS_PER_BLOCK = 4;
    
    int warp_id = threadIdx.x / 32;
    int lane_id = threadIdx.x % 32;
    int row = blockIdx.x * WARPS_PER_BLOCK + warp_id;
    int batch = blockIdx.y;
    
    if (row >= M) return;
    
    int64_t sfa_row_base = (row & 31) * sfa_s0 + ((row & 127) >> 5) * sfa_s1 + 
                           (row / 128) * sfa_s2 + batch * sfa_s5;
    int64_t sfb_batch_base = batch * sfb_s5;
    
    const uint8_t* A_row = A + row * a_s0 + batch * a_s2;
    const uint8_t* B_batch = B + batch * b_s2;
    
    half2 acc_h2 = __float2half2_rn(0.0f);
    int K_scales = K / 16;
    
    for (int scale_base = 0; scale_base < K_scales; scale_base += 32) {
        int k_scale = scale_base + lane_id;
        if (k_scale >= K_scales) break;
        
        int sfa_idx = sfa_row_base + (k_scale & 3) * sfa_s3 + (k_scale >> 2) * sfa_s4;
        int sfb_idx = sfb_batch_base + (k_scale & 3) * sfb_s3 + (k_scale >> 2) * sfb_s4;
        
        uint8_t sfa_val = SFA[sfa_idx];
        uint8_t sfb_val = SFB[sfb_idx];
        uint16_t sf_packed = (uint16_t(sfb_val) << 8) | uint16_t(sfa_val);
        uint32_t sf_f16x2 = cvt_f8x2_to_f16x2(sf_packed);
        half2 sf_h2 = *reinterpret_cast<half2*>(&sf_f16x2);
        half scale_h = __hmul(sf_h2.x, sf_h2.y);
        half2 scale_h2 = __half2half2(scale_h);
        
        int k_byte_start = k_scale * 8;
        uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte_start);
        uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte_start);
        
        uint32_t a_f16_0, a_f16_1, a_f16_2, a_f16_3, a_f16_4, a_f16_5, a_f16_6, a_f16_7;
        cvt_f4x8_to_f16x8(a_vec.x, a_f16_0, a_f16_1, a_f16_2, a_f16_3);
        cvt_f4x8_to_f16x8(a_vec.y, a_f16_4, a_f16_5, a_f16_6, a_f16_7);
        
        uint32_t b_f16_0, b_f16_1, b_f16_2, b_f16_3, b_f16_4, b_f16_5, b_f16_6, b_f16_7;
        cvt_f4x8_to_f16x8(b_vec.x, b_f16_0, b_f16_1, b_f16_2, b_f16_3);
        cvt_f4x8_to_f16x8(b_vec.y, b_f16_4, b_f16_5, b_f16_6, b_f16_7);
        
        half2 ab0 = __hmul2(*reinterpret_cast<half2*>(&a_f16_0), *reinterpret_cast<half2*>(&b_f16_0));
        half2 ab1 = __hmul2(*reinterpret_cast<half2*>(&a_f16_1), *reinterpret_cast<half2*>(&b_f16_1));
        half2 ab2 = __hmul2(*reinterpret_cast<half2*>(&a_f16_2), *reinterpret_cast<half2*>(&b_f16_2));
        half2 ab3 = __hmul2(*reinterpret_cast<half2*>(&a_f16_3), *reinterpret_cast<half2*>(&b_f16_3));
        half2 ab4 = __hmul2(*reinterpret_cast<half2*>(&a_f16_4), *reinterpret_cast<half2*>(&b_f16_4));
        half2 ab5 = __hmul2(*reinterpret_cast<half2*>(&a_f16_5), *reinterpret_cast<half2*>(&b_f16_5));
        half2 ab6 = __hmul2(*reinterpret_cast<half2*>(&a_f16_6), *reinterpret_cast<half2*>(&b_f16_6));
        half2 ab7 = __hmul2(*reinterpret_cast<half2*>(&a_f16_7), *reinterpret_cast<half2*>(&b_f16_7));
        
        half2 sum01 = __hadd2(ab0, ab1);
        half2 sum23 = __hadd2(ab2, ab3);
        half2 sum45 = __hadd2(ab4, ab5);
        half2 sum67 = __hadd2(ab6, ab7);
        half2 sum0123 = __hadd2(sum01, sum23);
        half2 sum4567 = __hadd2(sum45, sum67);
        half2 local_sum = __hadd2(sum0123, sum4567);
        
        acc_h2 = __hfma2(local_sum, scale_h2, acc_h2);
    }
    
    float acc = __half2float(acc_h2.x) + __half2float(acc_h2.y);
    acc = warp_reduce_sum(acc);
    
    if (lane_id == 0) {
        C[row * c_s0 + batch * c_s2] = __float2half(acc);
    }
}

torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
    int M = a.size(0), K = a.size(1) * 2, L = a.size(2);
    constexpr int WARPS_PER_BLOCK = 4;
    dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
    dim3 block(32 * WARPS_PER_BLOCK);
    
    gemv_v8<<<grid, block>>>(
        reinterpret_cast<const uint8_t*>(a.data_ptr()),
        reinterpret_cast<const uint8_t*>(b.data_ptr()),
        reinterpret_cast<const uint8_t*>(sfa.data_ptr()),
        reinterpret_cast<const uint8_t*>(sfb.data_ptr()),
        reinterpret_cast<half*>(c.data_ptr()),
        M, K, L, a.stride(0), a.stride(2), b.stride(2),
        sfa.stride(0), sfa.stride(1), sfa.stride(2), sfa.stride(3), sfa.stride(4), sfa.stride(5),
        sfb.stride(3), sfb.stride(4), sfb.stride(5), c.stride(0), c.stride(2)
    );
    return c;
}
"""

cpp_source = "torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);"

module = load_inline(
    name='gemv_v8',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['gemv_cuda'],
    extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17', '--generate-code=arch=compute_100a,code=sm_100a'],
    verbose=True
)

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, sfa_perm, sfb_perm, c = data
    return module.gemv_cuda(a, b, sfa_perm, sfb_perm, c)
scrolls · 154 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 107169.

⋯ 5 unchanged lines
#include <cuda_fp16.h>
#include <cuda_fp8.h>
- __device__ __forceinline__ float fp8_to_float(uint8_t x) {
- __nv_fp8_e4m3 v; v.__x = x;
- return float(v);
- }
-
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
⋯ 1 unchanged lines
return val;
}
- // Convert 8 FP4 values (packed in 1 uint32) to 8 FP16 values (in 2 uint32s)
- // Uses PTX cvt.rn.f16x2.e2m1x2 instruction
__device__ __forceinline__ void cvt_f4x8_to_f16x8(uint32_t src, uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3) {
asm volatile(
"{ .reg .b8 b0, b1, b2, b3; "
⋯ 7 unchanged lines
);
}
- // V5: Native PTX FP4->FP16 conversion
- __global__ void gemv_v5(
+ __device__ __forceinline__ uint32_t cvt_f8x2_to_f16x2(uint16_t src) {
+ uint32_t dst;
+ asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(dst) : "h"(src));
+ return dst;
+ }
+
+ __global__ void gemv_v8(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
half* __restrict__ C,
int M, int K, int L,
- int64_t a_s0, int64_t a_s2,
- int64_t b_s2,
+ int64_t a_s0, int64_t a_s2, int64_t b_s2,
int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
int64_t c_s0, int64_t c_s2
) {
- const int WARPS_PER_BLOCK = 8;
+ constexpr int WARPS_PER_BLOCK = 4;
int warp_id = threadIdx.x / 32;
int lane_id = threadIdx.x % 32;
⋯ 2 unchanged lines
if (row >= M) return;
- // Precompute base addresses
int64_t sfa_row_base = (row & 31) * sfa_s0 + ((row & 127) >> 5) * sfa_s1 +
(row / 128) * sfa_s2 + batch * sfa_s5;
int64_t sfb_batch_base = batch * sfb_s5;
⋯ 1 unchanged lines
const uint8_t* A_row = A + row * a_s0 + batch * a_s2;
const uint8_t* B_batch = B + batch * b_s2;
- float acc = 0.0f;
+ half2 acc_h2 = __float2half2_rn(0.0f);
int K_scales = K / 16;
- // Process 2 scale groups (32 elements) per iteration using 128-bit loads
for (int scale_base = 0; scale_base < K_scales; scale_base += 32) {
int k_scale = scale_base + lane_id;
if (k_scale >= K_scales) break;
- // Load scale factors
int sfa_idx = sfa_row_base + (k_scale & 3) * sfa_s3 + (k_scale >> 2) * sfa_s4;
int sfb_idx = sfb_batch_base + (k_scale & 3) * sfb_s3 + (k_scale >> 2) * sfb_s4;
- float scale = fp8_to_float(SFA[sfa_idx]) * fp8_to_float(SFB[sfb_idx]);
- // Load 8 bytes = 16 FP4 elements, but we only need 8 for one scale group
- // Actually, let's load 4 bytes = 8 FP4 elements = half a scale group
- // Wait - 16 elements per scale group, 8 bytes per scale group
- // Let's load 8 bytes with uint2, process 16 elements
+ uint8_t sfa_val = SFA[sfa_idx];
+ uint8_t sfb_val = SFB[sfb_idx];
+ uint16_t sf_packed = (uint16_t(sfb_val) << 8) | uint16_t(sfa_val);
+ uint32_t sf_f16x2 = cvt_f8x2_to_f16x2(sf_packed);
+ half2 sf_h2 = *reinterpret_cast<half2*>(&sf_f16x2);
+ half scale_h = __hmul(sf_h2.x, sf_h2.y);
+ half2 scale_h2 = __half2half2(scale_h);
int k_byte_start = k_scale * 8;
uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte_start);
uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte_start);
- // Convert first 4 bytes (8 FP4 values) of A
- uint32_t a_f16_0, a_f16_1, a_f16_2, a_f16_3;
+ uint32_t a_f16_0, a_f16_1, a_f16_2, a_f16_3, a_f16_4, a_f16_5, a_f16_6, a_f16_7;
cvt_f4x8_to_f16x8(a_vec.x, a_f16_0, a_f16_1, a_f16_2, a_f16_3);
-
- // Convert second 4 bytes (8 FP4 values) of A
- uint32_t a_f16_4, a_f16_5, a_f16_6, a_f16_7;
cvt_f4x8_to_f16x8(a_vec.y, a_f16_4, a_f16_5, a_f16_6, a_f16_7);
- // Convert first 4 bytes (8 FP4 values) of B
- uint32_t b_f16_0, b_f16_1, b_f16_2, b_f16_3;
+ uint32_t b_f16_0, b_f16_1, b_f16_2, b_f16_3, b_f16_4, b_f16_5, b_f16_6, b_f16_7;
cvt_f4x8_to_f16x8(b_vec.x, b_f16_0, b_f16_1, b_f16_2, b_f16_3);
-
- // Convert second 4 bytes (8 FP4 values) of B
- uint32_t b_f16_4, b_f16_5, b_f16_6, b_f16_7;
cvt_f4x8_to_f16x8(b_vec.y, b_f16_4, b_f16_5, b_f16_6, b_f16_7);
- // Compute dot product using half2 operations
- half2* a_h2_0 = reinterpret_cast<half2*>(&a_f16_0);
- half2* a_h2_1 = reinterpret_cast<half2*>(&a_f16_1);
- half2* a_h2_2 = reinterpret_cast<half2*>(&a_f16_2);
- half2* a_h2_3 = reinterpret_cast<half2*>(&a_f16_3);
- half2* a_h2_4 = reinterpret_cast<half2*>(&a_f16_4);
- half2* a_h2_5 = reinterpret_cast<half2*>(&a_f16_5);
- half2* a_h2_6 = reinterpret_cast<half2*>(&a_f16_6);
- half2* a_h2_7 = reinterpret_cast<half2*>(&a_f16_7);
+ half2 ab0 = __hmul2(*reinterpret_cast<half2*>(&a_f16_0), *reinterpret_cast<half2*>(&b_f16_0));
+ half2 ab1 = __hmul2(*reinterpret_cast<half2*>(&a_f16_1), *reinterpret_cast<half2*>(&b_f16_1));
+ half2 ab2 = __hmul2(*reinterpret_cast<half2*>(&a_f16_2), *reinterpret_cast<half2*>(&b_f16_2));
+ half2 ab3 = __hmul2(*reinterpret_cast<half2*>(&a_f16_3), *reinterpret_cast<half2*>(&b_f16_3));
+ half2 ab4 = __hmul2(*reinterpret_cast<half2*>(&a_f16_4), *reinterpret_cast<half2*>(&b_f16_4));
+ half2 ab5 = __hmul2(*reinterpret_cast<half2*>(&a_f16_5), *reinterpret_cast<half2*>(&b_f16_5));
+ half2 ab6 = __hmul2(*reinterpret_cast<half2*>(&a_f16_6), *reinterpret_cast<half2*>(&b_f16_6));
+ half2 ab7 = __hmul2(*reinterpret_cast<half2*>(&a_f16_7), *reinterpret_cast<half2*>(&b_f16_7));
- half2* b_h2_0 = reinterpret_cast<half2*>(&b_f16_0);
- half2* b_h2_1 = reinterpret_cast<half2*>(&b_f16_1);
- half2* b_h2_2 = reinterpret_cast<half2*>(&b_f16_2);
- half2* b_h2_3 = reinterpret_cast<half2*>(&b_f16_3);
- half2* b_h2_4 = reinterpret_cast<half2*>(&b_f16_4);
- half2* b_h2_5 = reinterpret_cast<half2*>(&b_f16_5);
- half2* b_h2_6 = reinterpret_cast<half2*>(&b_f16_6);
- half2* b_h2_7 = reinterpret_cast<half2*>(&b_f16_7);
-
- // Multiply and accumulate
- half2 prod0 = __hmul2(*a_h2_0, *b_h2_0);
- half2 prod1 = __hmul2(*a_h2_1, *b_h2_1);
- half2 prod2 = __hmul2(*a_h2_2, *b_h2_2);
- half2 prod3 = __hmul2(*a_h2_3, *b_h2_3);
- half2 prod4 = __hmul2(*a_h2_4, *b_h2_4);
- half2 prod5 = __hmul2(*a_h2_5, *b_h2_5);
- half2 prod6 = __hmul2(*a_h2_6, *b_h2_6);
- half2 prod7 = __hmul2(*a_h2_7, *b_h2_7);
-
- // Sum all products
- half2 sum01 = __hadd2(prod0, prod1);
- half2 sum23 = __hadd2(prod2, prod3);
- half2 sum45 = __hadd2(prod4, prod5);
- half2 sum67 = __hadd2(prod6, prod7);
-
+ half2 sum01 = __hadd2(ab0, ab1);
+ half2 sum23 = __hadd2(ab2, ab3);
+ half2 sum45 = __hadd2(ab4, ab5);
+ half2 sum67 = __hadd2(ab6, ab7);
half2 sum0123 = __hadd2(sum01, sum23);
half2 sum4567 = __hadd2(sum45, sum67);
+ half2 local_sum = __hadd2(sum0123, sum4567);
- half2 sum_all = __hadd2(sum0123, sum4567);
-
- float local_sum = __half2float(sum_all.x) + __half2float(sum_all.y);
-
- acc += local_sum * scale;
+ acc_h2 = __hfma2(local_sum, scale_h2, acc_h2);
}
+ float acc = __half2float(acc_h2.x) + __half2float(acc_h2.y);
acc = warp_reduce_sum(acc);
+
if (lane_id == 0) {
C[row * c_s0 + batch * c_s2] = __float2half(acc);
}
}
- torch::Tensor gemv_cuda(
- torch::Tensor a, torch::Tensor b,
- torch::Tensor sfa, torch::Tensor sfb,
- torch::Tensor c
- ) {
+ torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
int M = a.size(0), K = a.size(1) * 2, L = a.size(2);
-
- const int WARPS_PER_BLOCK = 8;
+ constexpr int WARPS_PER_BLOCK = 4;
dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
dim3 block(32 * WARPS_PER_BLOCK);
- gemv_v5<<<grid, block>>>(
+ gemv_v8<<<grid, block>>>(
reinterpret_cast<const uint8_t*>(a.data_ptr()),
reinterpret_cast<const uint8_t*>(b.data_ptr()),
reinterpret_cast<const uint8_t*>(sfa.data_ptr()),
reinterpret_cast<const uint8_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(c.data_ptr()),
- M, K, L,
- a.stride(0), a.stride(2),
- b.stride(2),
+ M, K, L, a.stride(0), a.stride(2), b.stride(2),
sfa.stride(0), sfa.stride(1), sfa.stride(2), sfa.stride(3), sfa.stride(4), sfa.stride(5),
- sfb.stride(3), sfb.stride(4), sfb.stride(5),
- c.stride(0), c.stride(2)
+ sfb.stride(3), sfb.stride(4), sfb.stride(5), c.stride(0), c.stride(2)
);
return c;
}
"""
- cpp_source = """
- torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
- """
+ cpp_source = "torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);"
module = load_inline(
- name='gemv_v5',
+ name='gemv_v8',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['gemv_cuda'],
scrolls · 223 diff lines total

Best evidence level for this revision: reported

JSON