Skip to content
KernelIndex
Search⌘K

submission 112996

shigao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

node_12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-112996?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
37.4µs
#204 of 678
2025-11-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:51d9b146a2795ef5858034cc6a1c353239553098b6e9625ced4e6ed9dc2cc1c6
license declaredunknown
license concludedunknown
authorsshigao
imported2026-08-15

Techniques

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

fp4int vec_count = K / 32; // int4 count for packed FP4 (16B per chunk)
fp8__device__ __forceinline__ __nv_fp8_e4m3 byte_to_fp8(uint8_t v) {
shared-memoryextern __shared__ uint8_t smem[];
vector-width = int4const int4* __restrict__ row_a_vec,

Kernel source

node_12.py411 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import TypeVar

input_t = TypeVar("input_t", bound=tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor])
output_t = TypeVar("output_t", bound=torch.Tensor)


cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <ATen/cuda/Exceptions.h>
#include <ATen/cuda/CUDAContext.h>
#include <sm_61_intrinsics.h>
#include <vector>

// Precomputed FP4 (e2m1) → int8 lookup table for DP4A.
__device__ __align__(16) int PACKED_FP4_DP4A[65536];
constexpr float kFp4Rescale = 1.0f / 16.0f;  // both operands scaled by 4

__device__ __forceinline__ int load_packed(uint16_t word) {
    return __ldg(&PACKED_FP4_DP4A[word]);
}

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

__device__ __forceinline__ __nv_fp8_e4m3 byte_to_fp8(uint8_t v) {
    union {
        uint8_t u;
        __nv_fp8_e4m3 fp;
    } cvt;
    cvt.u = v;
    return cvt.fp;
}

__device__ __forceinline__ float accumulate_chunk_single(
    int idx,
    const int4* __restrict__ row_a_vec,
    const int* __restrict__ row_b_packed,
    const uchar2* __restrict__ sfa_pairs,
    const float2* __restrict__ sfb_scales
) {
    int4 va = __ldg(row_a_vec + idx);
    const uint16_t* raw_a16 = reinterpret_cast<const uint16_t*>(&va);

    uchar2 sfa_pair = __ldg(sfa_pairs + idx);
    float2 sfb_scale = sfb_scales[idx];

    float scale_0 = static_cast<float>(byte_to_fp8(sfa_pair.x)) * sfb_scale.x;
    float scale_1 = static_cast<float>(byte_to_fp8(sfa_pair.y)) * sfb_scale.y;

    int base = idx * 8;

    int acc_int0 = 0;
    #pragma unroll
    for (int j = 0; j < 4; ++j) {
        int b_val = row_b_packed[base + j];
        acc_int0 = __dp4a(load_packed(raw_a16[j]), b_val, acc_int0);
    }

    int acc_int1 = 0;
    #pragma unroll
    for (int j = 4; j < 8; ++j) {
        int b_val = row_b_packed[base + j];
        acc_int1 = __dp4a(load_packed(raw_a16[j]), b_val, acc_int1);
    }

    float local0 = static_cast<float>(acc_int0) * kFp4Rescale;
    float local1 = static_cast<float>(acc_int1) * kFp4Rescale;
    return fmaf(local0, scale_0, local1 * scale_1);
}

__device__ __forceinline__ void accumulate_chunk_dual(
    int idx,
    const int4* __restrict__ row_a0_vec,
    const int4* __restrict__ row_a1_vec,
    const int* __restrict__ row_b_packed,
    const uchar2* __restrict__ sfa0_pairs,
    const uchar2* __restrict__ sfa1_pairs,
    const float2* __restrict__ sfb_scales,
    float& acc0,
    float& acc1
) {
    int4 va0 = __ldg(row_a0_vec + idx);
    int4 va1 = __ldg(row_a1_vec + idx);
    const uint16_t* raw_a16_0 = reinterpret_cast<const uint16_t*>(&va0);
    const uint16_t* raw_a16_1 = reinterpret_cast<const uint16_t*>(&va1);

    uchar2 sfa0_pair = __ldg(sfa0_pairs + idx);
    uchar2 sfa1_pair = __ldg(sfa1_pairs + idx);
    float2 sfb_scale = sfb_scales[idx];

    float scale00 = static_cast<float>(byte_to_fp8(sfa0_pair.x)) * sfb_scale.x;
    float scale01 = static_cast<float>(byte_to_fp8(sfa0_pair.y)) * sfb_scale.y;
    float scale10 = static_cast<float>(byte_to_fp8(sfa1_pair.x)) * sfb_scale.x;
    float scale11 = static_cast<float>(byte_to_fp8(sfa1_pair.y)) * sfb_scale.y;

    int base = idx * 8;

    int acc0_0 = 0, acc0_1 = 0;
    int acc1_0 = 0, acc1_1 = 0;

    #pragma unroll
    for (int j = 0; j < 4; ++j) {
        int b_val = row_b_packed[base + j];
        acc0_0 = __dp4a(load_packed(raw_a16_0[j]), b_val, acc0_0);
        acc1_0 = __dp4a(load_packed(raw_a16_1[j]), b_val, acc1_0);
    }

    #pragma unroll
    for (int j = 4; j < 8; ++j) {
        int b_val = row_b_packed[base + j];
        acc0_1 = __dp4a(load_packed(raw_a16_0[j]), b_val, acc0_1);
        acc1_1 = __dp4a(load_packed(raw_a16_1[j]), b_val, acc1_1);
    }

    float scaled0 = static_cast<float>(acc0_0) * kFp4Rescale;
    float scaled1 = static_cast<float>(acc0_1) * kFp4Rescale;
    acc0 = fmaf(scaled0, scale00, scaled1 * scale01) + acc0;

    float scaled2 = static_cast<float>(acc1_0) * kFp4Rescale;
    float scaled3 = static_cast<float>(acc1_1) * kFp4Rescale;
    acc1 = fmaf(scaled2, scale10, scaled3 * scale11) + acc1;
}

template<int WARPS_PER_BLOCK>
__global__ void nvfp4_gemv_shared_kernel(
    const uint8_t* __restrict__ A,          // [M, K/2, L]
    const uint8_t* __restrict__ B,          // [1, K/2, L] (only first row used)
    const __nv_fp8_e4m3* __restrict__ SFA,  // [M, K/16, L]
    const __nv_fp8_e4m3* __restrict__ SFB,  // [1, K/16, L]
    half* __restrict__ C,                   // [M, 1, L]
    int M, int K, int L,
    int64_t stride_a_m, int64_t stride_a_l,
    int64_t stride_b_l,
    int64_t stride_sfa_m, int64_t stride_sfa_l,
    int64_t stride_sfb_l,
    int64_t stride_c_m, int64_t stride_c_l
) {
    int tid = threadIdx.x;
    int warp_id = tid >> 5;
    int lane = tid & 31;
    int base_m = blockIdx.x * (WARPS_PER_BLOCK * 2);
    int l = blockIdx.y;

    if (base_m >= M || l >= L) {
        return;
    }

    extern __shared__ uint8_t smem[];
    int packed_words = K >> 2;   // K/4 packed int32 entries

    int b_packed_bytes = packed_words * sizeof(int);
    int b_packed_aligned = (b_packed_bytes + 15) & ~15;
    uint8_t* smem_b = smem;
    int* shared_b_packed = reinterpret_cast<int*>(smem_b);

    float2* shared_sfb_scale = reinterpret_cast<float2*>(smem_b + b_packed_aligned);

    // Cache packed B and SFB for this batch l once per block.
    const uint8_t* b_global = B + l * stride_b_l;
    const __nv_fp8_e4m3* sfb_global = SFB + l * stride_sfb_l;

    const uint16_t* b_words_global = reinterpret_cast<const uint16_t*>(b_global);
    for (int idx = tid; idx < packed_words; idx += blockDim.x) {
        uint16_t w = __ldg(b_words_global + idx);
        shared_b_packed[idx] = load_packed(w);
    }

    int sfb_pairs = K / 32; // two fp8 per pair
    const uchar2* sfb_vec_global = reinterpret_cast<const uchar2*>(sfb_global);
    for (int idx = tid; idx < sfb_pairs; idx += blockDim.x) {
        uchar2 v = sfb_vec_global[idx];
        float2 scale_pair;
        scale_pair.x = static_cast<float>(byte_to_fp8(v.x));
        scale_pair.y = static_cast<float>(byte_to_fp8(v.y));
        shared_sfb_scale[idx] = scale_pair;
    }

    __syncthreads();

    if (warp_id >= WARPS_PER_BLOCK) return;

    int m0 = base_m + warp_id * 2;
    int m1 = m0 + 1;
    bool has_row0 = m0 < M;
    bool has_row1 = m1 < M;
    if (!has_row0) return;

    const uint8_t* row_a0 = A + m0 * stride_a_m + l * stride_a_l;
    const __nv_fp8_e4m3* row_sfa0 = SFA + m0 * stride_sfa_m + l * stride_sfa_l;

    const int4* row_a0_vec = reinterpret_cast<const int4*>(row_a0);
    const uchar2* sfa0_pairs = reinterpret_cast<const uchar2*>(row_sfa0);

    const int* row_b_packed = shared_b_packed;
    const float2* sfb_scales = shared_sfb_scale;

    int vec_count = K / 32;   // int4 count for packed FP4 (16B per chunk)

    if (has_row1) {
        const uint8_t* row_a1 = A + m1 * stride_a_m + l * stride_a_l;
        const __nv_fp8_e4m3* row_sfa1 = SFA + m1 * stride_sfa_m + l * stride_sfa_l;
        const int4* row_a1_vec = reinterpret_cast<const int4*>(row_a1);
        const uchar2* sfa1_pairs = reinterpret_cast<const uchar2*>(row_sfa1);

        float acc_row0 = 0.0f;
        float acc_row1 = 0.0f;

        // Double-stripmine to reduce loop overhead while keeping ILP.
        for (int idx = lane; idx < vec_count; idx += 64) {
            accumulate_chunk_dual(idx, row_a0_vec, row_a1_vec, row_b_packed, sfa0_pairs, sfa1_pairs, sfb_scales, acc_row0, acc_row1);
            int idx_next = idx + 32;
            if (idx_next < vec_count) {
                accumulate_chunk_dual(idx_next, row_a0_vec, row_a1_vec, row_b_packed, sfa0_pairs, sfa1_pairs, sfb_scales, acc_row0, acc_row1);
            }
        }

        acc_row0 = warp_reduce_sum(acc_row0);
        acc_row1 = warp_reduce_sum(acc_row1);

        if (lane == 0) {
            int64_t out0 = static_cast<int64_t>(m0) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
            C[out0] = __float2half(acc_row0);
            int64_t out1 = static_cast<int64_t>(m1) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
            C[out1] = __float2half(acc_row1);
        }
    } else {
        float acc = 0.0f;

        for (int idx = lane; idx < vec_count; idx += 64) {
            acc += accumulate_chunk_single(idx, row_a0_vec, row_b_packed, sfa0_pairs, sfb_scales);
            int idx_next = idx + 32;
            if (idx_next < vec_count) {
                acc += accumulate_chunk_single(idx_next, row_a0_vec, row_b_packed, sfa0_pairs, sfb_scales);
            }
        }

        acc = warp_reduce_sum(acc);
        if (lane == 0) {
            int64_t out_idx = static_cast<int64_t>(m0) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
            C[out_idx] = __float2half(acc);
        }
    }
}

template<int WARPS_PER_BLOCK>
void launch_gemv(
    const uint8_t* A, const uint8_t* B,
    const __nv_fp8_e4m3* SFA, const __nv_fp8_e4m3* SFB,
    half* C,
    int M, int K, int L,
    int64_t stride_a_m, int64_t stride_a_l,
    int64_t stride_b_l,
    int64_t stride_sfa_m, int64_t stride_sfa_l,
    int64_t stride_sfb_l,
    int64_t stride_c_m, int64_t stride_c_l,
    int shared_bytes
) {
    dim3 block(32 * WARPS_PER_BLOCK);
    dim3 grid((M + (2 * WARPS_PER_BLOCK) - 1) / (2 * WARPS_PER_BLOCK), L);
    nvfp4_gemv_shared_kernel<WARPS_PER_BLOCK><<<grid, block, shared_bytes>>>(
        A, B, SFA, SFB, C,
        M, K, L,
        stride_a_m, stride_a_l,
        stride_b_l,
        stride_sfa_m, stride_sfa_l,
        stride_sfb_l,
        stride_c_m, stride_c_l
    );
}

void nvfp4_gemv(
    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; // logical elements
    int L = a.size(2);
    TORCH_CHECK(K % 64 == 0, "K must be divisible by 64");

    static bool lut_initialized = false;
    if (!lut_initialized) {
        static const int8_t E2M1_SCALED[16] = {
            0, 2, 4, 6, 8, 12, 16, 24,
            0, -2, -4, -6, -8, -12, -16, -24
        };
        std::vector<int> pack_lut(65536);
        for (int word = 0; word < 65536; ++word) {
            uint16_t v = static_cast<uint16_t>(word);
            int8_t e0 = E2M1_SCALED[v & 0x0F];
            int8_t e1 = E2M1_SCALED[(v >> 4) & 0x0F];
            int8_t e2 = E2M1_SCALED[(v >> 8) & 0x0F];
            int8_t e3 = E2M1_SCALED[(v >> 12) & 0x0F];
            uint32_t packed =
                static_cast<uint8_t>(e0) |
                (static_cast<uint8_t>(e1) << 8) |
                (static_cast<uint8_t>(e2) << 16) |
                (static_cast<uint8_t>(e3) << 24);
            pack_lut[word] = static_cast<int>(packed);
        }
        cudaMemcpyToSymbol(PACKED_FP4_DP4A, pack_lut.data(), pack_lut.size() * sizeof(int));
        lut_initialized = true;
    }

    int b_pack_bytes = K;       // K/4 int entries → K bytes
    int sfb_scale_bytes = (K >> 5) * static_cast<int>(sizeof(float2)); // K/32 pairs of float2 scales
    int shared_bytes = ((b_pack_bytes + 15) & ~15) + ((sfb_scale_bytes + 15) & ~15);

    const uint8_t* a_ptr = (const uint8_t*)a.data_ptr();
    const uint8_t* b_ptr = (const uint8_t*)b.data_ptr();
    const __nv_fp8_e4m3* sfa_ptr = (const __nv_fp8_e4m3*)sfa.data_ptr();
    const __nv_fp8_e4m3* sfb_ptr = (const __nv_fp8_e4m3*)sfb.data_ptr();
    half* c_ptr = (half*)c.data_ptr();

    // Warp selection tuned for dual-row processing.
    // Use 8 warps for medium/large K to keep occupancy healthy; 4 for very small.
    if (K <= 4096) {
        launch_gemv<4>(
            a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
            M, K, L,
            a.stride(0), a.stride(2),
            b.stride(2),
            sfa.stride(0), sfa.stride(2),
            sfb.stride(2),
            c.stride(0), c.stride(2),
            shared_bytes
        );
    } else {
        launch_gemv<8>(
            a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
            M, K, L,
            a.stride(0), a.stride(2),
            b.stride(2),
            sfa.stride(0), sfa.stride(2),
            sfb.stride(2),
            c.stride(0), c.stride(2),
            shared_bytes
        );
    }

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        printf("Kernel failed: %s\n", cudaGetErrorString(err));
    }
}
"""

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

nvfp4_gemv_module = load_inline(
    name='nvfp4_gemv',
    cpp_sources=[cpp_source],
    cuda_sources=[cuda_source],
    functions=['nvfp4_gemv'],
    verbose=True,
    extra_cflags=[
        '-std=c++17',
        "-O3",                      # CPU 代码最高优化
        "-march=native",            # 🚀 针对本机 CPU 指令集优化 (减少 launch overhead)
        "-fno-math-errno",          # 移除数学错误检查
        "-Wall",                    # 打开警告,防止低级错误
    ],
    extra_cuda_cflags=[
        "-O3",
        "--use_fast_math",
        '--extra-device-vectorization',
        "--restrict",
        "-std=c++17",
        "--ptxas-options=-O3",
        "--expt-relaxed-constexpr",
        "-arch=sm_100a",
        '-Xptxas', '-v',
        '-lineinfo',
        '-U__CUDA_NO_HALF_OPERATORS__',  # Enable half operators
        '-U__CUDA_NO_HALF_CONVERSIONS__',  # Enable half conversions
    ],
)

def custom_kernel(data: input_t) -> output_t:
    a_ref, b_ref, sfa, sfb, _, _, c_ref = data

    nvfp4_gemv_module.nvfp4_gemv(
        a_ref,
        b_ref,
        sfa,
        sfb,
        c_ref
    )

    return c_ref
scrolls · 411 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