Skip to content
KernelIndex
Search⌘K

submission 107251

Rudrashis · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_rudrashis.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-107251?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
31.8µs
#168 of 678
2025-11-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0e149a5a37f74a2a1d79def016c6df728edfbedc4a32fd557d66f77b46d33e22
license declaredunknown
license concludedunknown
authorsRudrashis
imported2026-08-15

Techniques

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

async-copyasm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
fp8__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
shared-memory__shared__ float smem_acc[BLOCK_SIZE];
vector-width = float2float2 f0 = __half22float2(acc_h2_0);

Kernel source

sub_rudrashis.py417 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t


cpp_source = """
#include <torch/extension.h>
#include <cuda_runtime.h>

torch::Tensor batched_scaled_gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""

cuda_source = """
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>

#define BLOCK_SIZE 32
#define K_TILE 2560
#define SCALES_PER_TILE (K_TILE / 16)
#define BYTES_PER_TILE (K_TILE / 2)
#define NUM_BUFFERS 2

__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
    __half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
        static_cast<__nv_fp4x2_storage_t>(byte),
        __NV_E2M1
    );
    return *reinterpret_cast<__half2*>(&raw);
}

__device__ __forceinline__ float decode_fp8(int8_t byte) {
    __nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
    __half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
    return __half2float(__ushort_as_half(raw.x));
}

__device__ __forceinline__ __half2 dot_scaled_4bytes(
    uint32_t a4,
    uint32_t b4,
    __half2 scale_h2
) {
    // Extract bytes using bfe.u32 for clearer semantics and same performance as prmt
    uint32_t b_byte0, b_byte1, b_byte2, b_byte3;
    uint32_t a_byte0, a_byte1, a_byte2, a_byte3;

    // Extract 4 bytes from b4 using bit field extract
    asm("bfe.u32 %0, %1, 0, 8;"  : "=r"(b_byte0) : "r"(b4));
    asm("bfe.u32 %0, %1, 8, 8;"  : "=r"(b_byte1) : "r"(b4));
    asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b_byte2) : "r"(b4));
    asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b_byte3) : "r"(b4));

    // Extract 4 bytes from a4
    asm("bfe.u32 %0, %1, 0, 8;"  : "=r"(a_byte0) : "r"(a4));
    asm("bfe.u32 %0, %1, 8, 8;"  : "=r"(a_byte1) : "r"(a4));
    asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a_byte2) : "r"(a4));
    asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a_byte3) : "r"(a4));

    // Inline b_scaled computation to minimize register live ranges
    __half2 acc = __hmul2(decode_fp4x2(a_byte0), __hmul2(decode_fp4x2(b_byte0), scale_h2));
    acc = __hfma2(decode_fp4x2(a_byte1), __hmul2(decode_fp4x2(b_byte1), scale_h2), acc);
    acc = __hfma2(decode_fp4x2(a_byte2), __hmul2(decode_fp4x2(b_byte2), scale_h2), acc);
    acc = __hfma2(decode_fp4x2(a_byte3), __hmul2(decode_fp4x2(b_byte3), scale_h2), acc);

    return acc;
}

__device__ __forceinline__ float compute_tile(
    const uint8_t* sh_a,
    const uint8_t* sh_b,
    const uint8_t* sh_sfa,
    const uint8_t* sh_sfb,
    int tid
) {
    float acc = 0.0f;

#pragma unroll 8
    for (int sf = tid; sf < SCALES_PER_TILE; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
                      decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;  // sf * 8 using bit shift

        uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
        uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
        uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
        uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        // Combine and accumulate in one step
        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

__device__ __forceinline__ float compute_remainder(
    const uint8_t* row_a,
    const uint8_t* batch_b,
    const uint8_t* row_sfa,
    const uint8_t* batch_sfb,
    int remainder_sf_start,
    int K_sf,
    int tid
) {
    float acc = 0.0f;

#pragma unroll 4
    for (int sf = remainder_sf_start + tid; sf < K_sf; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(__ldg(&row_sfa[sf]))) *
                      decode_fp8(static_cast<int8_t>(__ldg(&batch_sfb[sf])));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;  // sf * 8 using bit shift

        uint32_t a4_0 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base]));
        uint32_t b4_0 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base]));
        uint32_t a4_1 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base + 4]));
        uint32_t b4_1 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base + 4]));

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        // Combine and accumulate in one step
        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

__device__ __forceinline__ float compute_remainder_smem(
    const uint8_t* sh_a,
    const uint8_t* sh_b,
    const uint8_t* sh_sfa,
    const uint8_t* sh_sfb,
    int remainder_scales,
    int tid
) {
    float acc = 0.0f;

#pragma unroll 4
    for (int sf = tid; sf < remainder_scales; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
                      decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;  // sf * 8 using bit shift

        uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
        uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
        uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
        uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        // Combine and accumulate in one step
        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

__global__ void gemv_nvfp4_kernel(
    const int8_t* __restrict__ a,
    const int8_t* __restrict__ b,
    const int8_t* __restrict__ sfa,
    const int8_t* __restrict__ sfb,
    half* __restrict__ c,
    int M, int K, int L,
    int N_rows
) {
    int m = blockIdx.x;
    int l = blockIdx.y;
    int tid = threadIdx.x;

    if (m >= M || l >= L) return;

    const int K_sf = K / 16;
    const int K_half = K / 2;
    const size_t batch_stride_a = static_cast<size_t>(M) * K_half;
    const size_t batch_stride_b = static_cast<size_t>(N_rows) * K_half;
    const size_t batch_stride_sfa = static_cast<size_t>(M) * K_sf;
    const size_t batch_stride_sfb = static_cast<size_t>(N_rows) * K_sf;

    const uint8_t* base_a = reinterpret_cast<const uint8_t*>(a);
    const uint8_t* base_b = reinterpret_cast<const uint8_t*>(b);
    const uint8_t* base_sfa = reinterpret_cast<const uint8_t*>(sfa);
    const uint8_t* base_sfb = reinterpret_cast<const uint8_t*>(sfb);

    const uint8_t* batch_a = base_a + l * batch_stride_a;
    const uint8_t* batch_b = base_b + l * batch_stride_b;
    const uint8_t* batch_sfa = base_sfa + l * batch_stride_sfa;
    const uint8_t* batch_sfb = base_sfb + l * batch_stride_sfb;

    const uint8_t* row_a = batch_a + static_cast<size_t>(m) * K_half;
    const uint8_t* row_sfa = batch_sfa + static_cast<size_t>(m) * K_sf;

    __shared__ float smem_acc[BLOCK_SIZE];
    __shared__ uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];
    __shared__ uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
    __shared__ uint8_t sh_sfa[NUM_BUFFERS][SCALES_PER_TILE];
    __shared__ uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];

    float acc = 0.0f;
    int tile_count = K / K_TILE;
    int remainder_start = tile_count * K_TILE;

#define ASYNC_COPY_16(dst, src) \
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))

#define ASYNC_COPY_4(dst, src) \
    asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))

    auto issue_tile_async = [&](int b_idx, int tile_idx) {
        const int base_k = tile_idx * K_TILE;
        const int base_byte = base_k >> 1;  // divide by 2 using shift
        const int base_sf = base_k >> 4;    // divide by 16 using shift
        const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
        const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
        const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
        const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);

#pragma unroll 2
        for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
            ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
            ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
        }
#pragma unroll
        for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
            ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
            ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
        }
        asm volatile("cp.async.commit_group;");
    };

    auto issue_remainder_async = [&](int b_idx, int rem_sf_start, int total_K_sf) {
        const int base_byte = rem_sf_start << 3;  // rem_sf_start * 16 / 2
        const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
        const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
        const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
        const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);

        int remainder_bytes = ((total_K_sf - rem_sf_start) << 3);  // * 8
        int remainder_scales = total_K_sf - rem_sf_start;

        // Copy remainder data - adapt stride to remainder size
        for (int i = tid * 16; i < remainder_bytes; i += BLOCK_SIZE * 16) {
            ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
            ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
        }
        for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {
            ASYNC_COPY_4(sh_sfa_base + i, row_sfa + rem_sf_start + i);
            ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);
        }
        asm volatile("cp.async.commit_group;");
    };

    int remainder_sf_start = remainder_start / 16;
    bool has_remainder = (remainder_sf_start < K_sf);
    int buf = 0;  // Declare buf outside so it's accessible in remainder section

    if (tile_count > 0) {
        issue_tile_async(0, 0);
        asm volatile("cp.async.wait_group 0;");
        __syncthreads();

        for (int tile = 0; tile < tile_count; ++tile) {
            if (tile + 1 < tile_count) {
                // Prefetch next tile
                issue_tile_async(buf ^ 1, tile + 1);
            } else if (has_remainder) {
                // On last tile: prefetch remainder into unused buffer
                issue_remainder_async(buf ^ 1, remainder_sf_start, K_sf);
            }

            acc += compute_tile(
                sh_a[buf],
                sh_b[buf],
                sh_sfa[buf],
                sh_sfb[buf],
                tid
            );

            if (tile + 1 < tile_count || has_remainder) {
                asm volatile("cp.async.wait_group 0;");
                __syncthreads();
                buf ^= 1;
            }
        }
    }

    if (has_remainder) {
        if (tile_count > 0) {
            // Remainder was prefetched during last tile - use from shared memory
            int remainder_scales = K_sf - remainder_sf_start;
            acc += compute_remainder_smem(
                sh_a[buf],
                sh_b[buf],
                sh_sfa[buf],
                sh_sfb[buf],
                remainder_scales,
                tid
            );
        } else {
            // No tiles processed - load remainder directly from global memory
            acc += compute_remainder(
                row_a,
                batch_b,
                row_sfa,
                batch_sfb,
                remainder_sf_start,
                K_sf,
                tid
            );
        }
    }

#undef ASYNC_COPY_16
#undef ASYNC_COPY_4

    float warp_sum = acc;
    // Fully unrolled warp reduction for lower latency
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 16);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 8);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 4);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 2);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 1);

    const int warp_id = tid >> 5;
    const int lane = tid & 31;
    if (lane == 0) smem_acc[warp_id] = warp_sum;

    __syncthreads();

    if (warp_id == 0) {
        float block_sum = (lane < (BLOCK_SIZE >> 5)) ? smem_acc[lane] : 0.0f;
        // Fully unrolled block reduction
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 16);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 8);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 4);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 2);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 1);

        if (lane == 0) {
            size_t c_idx = static_cast<size_t>(m) + static_cast<size_t>(l) * M;
            c[c_idx] = __float2half(block_sum);
        }
    }
}

torch::Tensor batched_scaled_gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
    int M = a.size(0);
    int K = a.size(1) * 2;
    int L = a.size(2);
    int N_rows = b.size(0);

    dim3 grid(M, L);
    dim3 block(BLOCK_SIZE);

    auto* a_ptr = reinterpret_cast<const int8_t*>(a.data_ptr());
    auto* b_ptr = reinterpret_cast<const int8_t*>(b.data_ptr());
    auto* sfa_ptr = reinterpret_cast<const int8_t*>(sfa.data_ptr());
    auto* sfb_ptr = reinterpret_cast<const int8_t*>(sfb.data_ptr());
    auto* c_ptr = reinterpret_cast<half*>(c.data_ptr());

    gemv_nvfp4_kernel<<<grid, block>>>(
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
        M, K, L,
        N_rows
    );

    return c;
}

"""

module = load_inline(
    name='batched_scaled_gemv',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['batched_scaled_gemv_cuda'],
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '-std=c++17',
        '-gencode=arch=compute_100a,code=sm_100a'
    ],
    with_cuda=True,
    verbose=False
)

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa_ref, sfb_ref, _, _, c = data
    device = a.device
    return module.batched_scaled_gemv_cuda(
        a,
        b,
        sfa_ref,
        sfb_ref,
        c
    )
scrolls · 417 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