Skip to content
KernelIndex
Search⌘K

submission 105051

yue · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:01d5e4aeb9ff132e0020cf720976c12905df0418d11fee721ad1b0ded3949bb0
license declaredunknown
license concludedunknown
authorsyue
imported2026-08-15

Techniques

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

fp4- Optimized PTX decode+FMA for FP4 x8 (reduced instruction count)
shared-memory__shared__ __half smem_results[ROWS_PER_BLOCK];
tile-k = 64constexpr int TILE_K = 64; // Each thread handles 64 elements
vector-width = float4float4 A_data[2]; // 2x16 bytes = 32 bytes

Kernel source

submit_v1.py279 lines
import torch
import sys
from torch.utils.cpp_extension import load_inline
from typing import Tuple
from task import input_t, output_t

gemv_cuda_src = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda_pipeline.h>
#include <cstdint>

__device__ __forceinline__ float decode_mul_accumulate_fp4x8(
    const uint32_t a_packed, 
    const uint32_t b_packed, 
    float acc)
{
    float result;
    
    asm volatile (
        "{"
        "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
        "  .reg .b32 %%a<4>, %%b<4>;\\n"
        "  .reg .b32 %%p0, %%p1;\\n"
        "  .reg .f16 %%h0, %%h1;\\n"
        "  .reg .f32 %%f0, %%f1;\\n"
        "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
        "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
        "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
        "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
        "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
        "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
        "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
        "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
        "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
        "  cvt.f32.f16 %%f0, %%h0;\\n"
        "  cvt.f32.f16 %%f1, %%h1;\\n"
        "  add.f32 %%f0, %%f0, %%f1;\\n"
        "  add.f32 %0, %%f0, %3;\\n"
        "}"
        : "=f"(result)
        : "r"(a_packed), "r"(b_packed), "f"(acc)
    );
    
    return result;
}

__device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
{
    uint32_t out_low, out_high;
    
    asm volatile (
        "{"
        "  .reg .b16 %%low, %%high;\\n"
        "  mov.b32 {%%low, %%high}, %2;\\n"
        "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
        "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
        "}"
        : "=r"(out_low), "=r"(out_high)
        : "r"(packed)
    );
    
    __half2 h_low = *reinterpret_cast<const __half2*>(&out_low);
    __half2 h_high = *reinterpret_cast<const __half2*>(&out_high);
    h0 = h_low.x;
    h1 = h_low.y;
    h2 = h_high.x;
    h3 = h_high.y;
}

extern "C" __global__
void block_scaled_gemv_fp4_fp8_fp16_vectorized(
    const uint8_t* __restrict__ A, // [l, m, k//2]
    const uint8_t* __restrict__ B, // [l, 128, k//2]
    const uint8_t* __restrict__ SFA, // [l, m, k//16]
    const uint8_t* __restrict__ SFB, // [l, 128, k//16]
    __half* __restrict__ C, // [l, m, 1]
    int M, int K, int L)
{
    constexpr int ROWS_PER_BLOCK = 32;
    constexpr int THREADS_PER_ROW = 16; // 16 threads collaborate on each row
    constexpr int TILE_K = 64; // Each thread handles 64 elements
    constexpr int SF_VEC_SIZE = 16;

    const int tidx = threadIdx.x; // 0-15: which K-segment for this row
    const int tidy = threadIdx.y; // 0-31: which row
    const int block_row = blockIdx.x * ROWS_PER_BLOCK;
    const int batch_idx = blockIdx.z;
    const int global_row = block_row + tidy;

    // Shared memory for coalesced writes
    __shared__ __half smem_results[ROWS_PER_BLOCK];

    if (global_row >= M) return;

    float local_sum = 0.0f;
    const int num_k_tiles = K / TILE_K;

    const uint8_t* B_base = B + batch_idx * (128 * K / 2);
    const uint8_t* SFB_base = SFB + batch_idx * (128 * K / 16);
    const uint8_t* A_row = A + batch_idx * (M * K / 2) + global_row * (K / 2);
    const uint8_t* SFA_row = SFA + batch_idx * (M * K / 16) + global_row * (K / 16);

    // K-PARALLELISM: Each thread handles multiple K-tiles, striding by THREADS_PER_ROW
    #pragma unroll 1
    for (int tile = tidx; tile < num_k_tiles; tile += THREADS_PER_ROW) {
        const int k_offset = tile * TILE_K;

        // Vectorized loads
        float4 A_data[2]; // 2x16 bytes = 32 bytes
        float4 B_data[2];

        const float4* A_ptr = reinterpret_cast<const float4*>(A_row + k_offset / 2);
        const float4* B_ptr = reinterpret_cast<const float4*>(B_base + k_offset / 2);

        A_data[0] = __ldg(A_ptr);
        A_data[1] = __ldg(A_ptr + 1);
        B_data[0] = __ldg(B_ptr);
        B_data[1] = __ldg(B_ptr + 1);

        uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);
        uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);

        // Load scale factors (4 bytes per thread)
        uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + k_offset / 16));
        uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + k_offset / 16));

        // Vectorized decode for all 4 scales at once
        __half sfa_scales[4], sfb_scales[4];
        decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);
        decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);

        // COMPUTE with vectorized decode using PTX asm
        constexpr int NUM_SF_BLOCKS = TILE_K / SF_VEC_SIZE;
        
        #pragma unroll
        for (int sf_block = 0; sf_block < NUM_SF_BLOCKS; sf_block++) {
            const int base_k = sf_block * SF_VEC_SIZE;

            // Use pre-decoded scales
            __half sfa = sfa_scales[sf_block];
            __half sfb = sfb_scales[sf_block];
            __half scale = __hmul(sfa, sfb);
            float block_sum = 0.0f;

            // Process 16 elements (8 bytes packed) using optimized decode+FMA
            #pragma unroll
            for (int sub = 0; sub < 2; sub++) {
                const int byte_idx = base_k / 2 + sub * 4;

                const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);
                const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);

                // Optimized: decode and accumulate in fewer instructions
                block_sum = decode_mul_accumulate_fp4x8(a_packed, b_packed, block_sum);
            }

            local_sum = __fmaf_rn(__half2float(scale), block_sum, local_sum);
        }
    }

    // Warp-level reduction across K dimension (16 threads per row)
    constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads
    
    #pragma unroll
    for (int offset = 8; offset > 0; offset >>= 1) {
        local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);
    }

    // First thread in each row writes to shared memory
    if (tidx == 0) {
        smem_results[tidy] = __float2half(local_sum);
    }
    __syncthreads();

    // Coalesced write
    const int linear_tid = tidx + tidy * THREADS_PER_ROW;
    if (linear_tid < ROWS_PER_BLOCK) {
        const int write_row = block_row + linear_tid;
        if (write_row < M) {
            C[batch_idx * M + write_row] = smem_results[linear_tid];
        }
    }
}

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L)
{
    TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
    TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
    TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");
    TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");
    TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");

    constexpr int ROWS_PER_BLOCK = 32;
    constexpr int THREADS_PER_ROW = 16;

    const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);
    const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);

    block_scaled_gemv_fp4_fp8_fp16_vectorized<<<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);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess)
        throw std::runtime_error(cudaGetErrorString(err));

    return C;
}
"""

gemv_cpp_src = """
#include <torch/extension.h>

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L);
"""

_gemm_module = load_inline(
    name="block_scaled_gemv_vectorized_v1",
    cpp_sources=gemv_cpp_src,
    cuda_sources=gemv_cuda_src,
    functions=["gemv_fp4_fp8_fp16"],
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '-std=c++17',
        '--expt-relaxed-constexpr',
        '--maxrregcount=256',
        '-gencode=arch=compute_100a,code=sm_100a',
    ],
    verbose=True,
)

def custom_kernel(data: input_t) -> output_t:
    """
    Optimized K-parallel with PTX FMA instructions:
    - K-parallelism (one thread = multiple tiles)
    - Vectorized loads using __ldg and float4
    - Warp shuffle reduction
    - Optimized PTX decode+FMA for FP4 x8 (reduced instruction count)
    - PTX FMA for scale multiplication and accumulation
    """
    a, b, sfa, sfb, _, _, c = data
    m, k_packed, l = a.shape
    k = k_packed * 2

    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)
    sfa_uint8 = sfa.view(torch.uint8)
    sfb_uint8 = sfb.view(torch.uint8)

    _gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)

    return c

scrolls · 279 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 104932.

⋯ 9 unchanged lines
#include <cuda_pipeline.h>
#include <cstdint>
- // Vectorized decode for FP4 x8 - FIXED VERSION
- // cvt.rn.f16x2.e2m1x2 takes .b8 input and produces .b32 output (f16x2)
- __device__ __forceinline__ void decode_fp4x8_e2m1_half4(const uint32_t packed, __half2& h0, __half2& h1, __half2& h2, __half2& h3)
+ __device__ __forceinline__ float decode_mul_accumulate_fp4x8(
+ const uint32_t a_packed,
+ const uint32_t b_packed,
+ float acc)
{
- uint32_t out0, out1, out2, out3;
+ float result;
- // PTX allows wider operands for .b8 instruction types
- // Input is uint32 containing 4 bytes, output is 4x .b32 registers
asm volatile (
"{"
- " .reg .b8 %%b0, %%b1, %%b2, %%b3;\\n"
- " mov.b32 {%%b0, %%b1, %%b2, %%b3}, %4;\\n"
- " cvt.rn.f16x2.e2m1x2 %0, %%b0;\\n"
- " cvt.rn.f16x2.e2m1x2 %1, %%b1;\\n"
- " cvt.rn.f16x2.e2m1x2 %2, %%b2;\\n"
- " cvt.rn.f16x2.e2m1x2 %3, %%b3;\\n"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %%f0, %%f0, %%f1;\\n"
+ " add.f32 %0, %%f0, %3;\\n"
"}"
- : "=r"(out0), "=r"(out1), "=r"(out2), "=r"(out3)
- : "r"(packed)
+ : "=f"(result)
+ : "r"(a_packed), "r"(b_packed), "f"(acc)
);
- h0 = *reinterpret_cast<const __half2*>(&out0);
- h1 = *reinterpret_cast<const __half2*>(&out1);
- h2 = *reinterpret_cast<const __half2*>(&out2);
- h3 = *reinterpret_cast<const __half2*>(&out3);
+ return result;
}
- // Vectorized decode for FP8 x4 - FIXED VERSION
- // cvt.rn.f16x2.e4m3x2 takes .b16 input and produces .b32 output (f16x2)
__device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
{
uint32_t out_low, out_high;
- // Input is uint32 containing 2x16-bit values
asm volatile (
"{"
" .reg .b16 %%low, %%high;\\n"
⋯ 88 unchanged lines
__half scale = __hmul(sfa, sfb);
float block_sum = 0.0f;
- // Process 16 elements (8 bytes packed) using two x8 decodes
+ // Process 16 elements (8 bytes packed) using optimized decode+FMA
#pragma unroll
for (int sub = 0; sub < 2; sub++) {
const int byte_idx = base_k / 2 + sub * 4;
⋯ 1 unchanged lines
const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);
const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);
- __half2 a0, a1, a2, a3;
- __half2 b0, b1, b2, b3;
- decode_fp4x8_e2m1_half4(a_packed, a0, a1, a2, a3);
- decode_fp4x8_e2m1_half4(b_packed, b0, b1, b2, b3);
-
- // Vectorized multiplies
- __half2 p0 = __hmul2(a0, b0);
- __half2 p1 = __hmul2(a1, b1);
- __half2 p2 = __hmul2(a2, b2);
- __half2 p3 = __hmul2(a3, b3);
-
- // Convert to float and accumulate
- float2 fp0 = __half22float2(p0);
- float2 fp1 = __half22float2(p1);
- float2 fp2 = __half22float2(p2);
- float2 fp3 = __half22float2(p3);
-
- block_sum += fp0.x + fp0.y;
- block_sum += fp1.x + fp1.y;
- block_sum += fp2.x + fp2.y;
- block_sum += fp3.x + fp3.y;
+ // Optimized: decode and accumulate in fewer instructions
+ block_sum = decode_mul_accumulate_fp4x8(a_packed, b_packed, block_sum);
}
- local_sum += __half2float(scale) * block_sum;
+ local_sum = __fmaf_rn(__half2float(scale), block_sum, local_sum);
}
}
⋯ 70 unchanged lines
"""
_gemm_module = load_inline(
- name="block_scaled_gemv_vectorized",
+ name="block_scaled_gemv_vectorized_v1",
cpp_sources=gemv_cpp_src,
cuda_sources=gemv_cuda_src,
functions=["gemv_fp4_fp8_fp16"],
⋯ 10 unchanged lines
def custom_kernel(data: input_t) -> output_t:
"""
- K-parallel with explicit vectorized loads and warp-level reduction:
- - Keeps original K-parallelism (one thread = multiple tiles)
- - Uses __ldg and float4 for explicit vectorized loads
+ Optimized K-parallel with PTX FMA instructions:
+ - K-parallelism (one thread = multiple tiles)
+ - Vectorized loads using __ldg and float4
- Warp shuffle reduction
- - Vectorized decodes using PTX asm for FP4 x8 and FP8 x4
+ - Optimized PTX decode+FMA for FP4 x8 (reduced instruction count)
+ - PTX FMA for scale multiplication and accumulation
"""
a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
⋯ 6 unchanged lines
_gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)
- return c
No newline at end of file
+ return c
+
scrolls · 151 diff lines total

Best evidence level for this revision: reported

JSON