Skip to content
KernelIndex
Search⌘K

submission 107169

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:54bf3fa35db9b9386de3f81793126c7c95083541653c68b7da3eb65413658714
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15

Techniques

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

fp8__nv_fp8_e4m3 v; v.__x = x;
vector-width = half2half2* a_h2_0 = reinterpret_cast<half2*>(&a_f16_0);

Kernel source

kmajor.py199 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 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)
        val += __shfl_xor_sync(0xffffffff, val, offset);
    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; "
        "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)
    );
}

// V5: Native PTX FP4->FP16 conversion
__global__ void gemv_v5(
    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
) {
    const int WARPS_PER_BLOCK = 8;
    
    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;
    
    // 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;
    
    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;
    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
        
        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;
        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;
        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* 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 sum0123 = __hadd2(sum01, sum23);
        half2 sum4567 = __hadd2(sum45, sum67);
        
        half2 sum_all = __hadd2(sum0123, sum4567);
        
        float local_sum = __half2float(sum_all.x) + __half2float(sum_all.y);
        
        acc += local_sum * scale;
    }
    
    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);
    
    const int WARPS_PER_BLOCK = 8;
    dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
    dim3 block(32 * WARPS_PER_BLOCK);
    
    gemv_v5<<<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_v5',
    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 · 199 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 70699.

import torch
from torch.utils.cpp_extension import load_inline
+ from task import input_t, output_t
- nvfp4_gemv_cuda = """
+ cuda_source = """
#include <cuda_fp16.h>
- #include <cuda_runtime.h>
+ #include <cuda_fp8.h>
- __constant__ float fp4_lut[16] = {
- 0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
- -0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
- };
+ __device__ __forceinline__ float fp8_to_float(uint8_t x) {
+ __nv_fp8_e4m3 v; v.__x = x;
+ return float(v);
+ }
- __constant__ float fp8_e4m3_lut[256];
-
- __device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
- return fp4_lut[(idx == 0) ? (packed_val & 0x0F) : (packed_val >> 4)];
+ __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__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
- return fp8_e4m3_lut[fp8_bits];
+ // 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; "
+ "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)
+ );
}
- template<int ROWS_PER_WARP, int VEC_SIZE>
- __global__ void nvfp4_gemv_kernel(
- const unsigned char* __restrict__ a,
- const unsigned char* __restrict__ b,
- const unsigned char* __restrict__ sfa,
- const unsigned char* __restrict__ sfb,
- __half* __restrict__ c,
- int M, int K, int L, int B_rows,
- int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,
- int64_t a_s0, int64_t a_s1, int64_t a_s2,
- int64_t b_s0, int64_t b_s1, int64_t b_s2,
+ // V5: Native PTX FP4->FP16 conversion
+ __global__ void gemv_v5(
+ 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_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
- int64_t c_s0, int64_t c_s1, int64_t c_s2
+ int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
+ int64_t c_s0, int64_t c_s2
) {
- extern __shared__ unsigned char smem[];
- unsigned char* b_shared = smem;
- unsigned char* sfb_shared = smem + (K / 2) + 4;
+ const int WARPS_PER_BLOCK = 8;
- const int warp_id = threadIdx.y;
- const int lane_id = threadIdx.x;
- const int l = blockIdx.y;
- const int tid = threadIdx.x + threadIdx.y * blockDim.x;
+ 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;
- const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;
+ if (row >= M) return;
- const int K_bytes = K / 2;
- const int K_blocks = K / 16;
- const int b_m = 0;
+ // 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;
- constexpr int B_VEC_SIZE = 16;
- for (int byte_idx = tid * B_VEC_SIZE; byte_idx < K_bytes; byte_idx += blockDim.x * blockDim.y * B_VEC_SIZE) {
- if (byte_idx + B_VEC_SIZE <= K_bytes && b_s1 == 1) {
- const int64_t b_offset_base = byte_idx * b_s1 + l * b_s2;
- *reinterpret_cast<uint4*>(&b_shared[byte_idx]) =
- __ldg(reinterpret_cast<const uint4*>(&b[b_offset_base]));
- } else {
- for (int i = 0; i < B_VEC_SIZE && byte_idx + i < K_bytes; ++i) {
- b_shared[byte_idx + i] = __ldg(&b[(byte_idx + i) * b_s1 + l * b_s2]);
- }
- }
- }
+ const uint8_t* A_row = A + row * a_s0 + batch * a_s2;
+ const uint8_t* B_batch = B + batch * b_s2;
- const int64_t sfb_l_offset = l * sfb_s5;
- for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {
- const int kk = idx / 4;
- const int kk4 = idx % 4;
- const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;
- sfb_shared[idx] = __ldg(&sfb[sfb_offset]);
- }
- __syncthreads();
+ float acc = 0.0f;
+ int K_scales = K / 16;
- float thread_acc[ROWS_PER_WARP];
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) thread_acc[r] = 0.0f;
-
- int64_t sfa_m_parts[ROWS_PER_WARP];
- bool row_valid[ROWS_PER_WARP];
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- const int m = base_m + r;
- row_valid[r] = (m < M && l < L);
- if (row_valid[r]) {
- const int mm = m / 128;
- const int mm32 = m % 32;
- const int mm4 = (m % 128) / 32;
- sfa_m_parts[r] = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;
- }
- }
-
- const int total_vec_iters = K_bytes / (warpSize * VEC_SIZE);
- const int remainder_start = total_vec_iters * warpSize * VEC_SIZE;
-
- for (int iter = 0; iter < total_vec_iters; ++iter) {
- const int k_byte_base = iter * warpSize * VEC_SIZE + lane_id * VEC_SIZE;
+ // 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;
- // Process in chunks to reduce register pressure
- constexpr int CHUNK_SIZE = VEC_SIZE / 2;
+ // 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]);
- for (int chunk = 0; chunk < 2; ++chunk) {
- const int chunk_offset = chunk * CHUNK_SIZE;
- const int block_0 = (k_byte_base + chunk_offset) >> 3;
- const int block_1 = (k_byte_base + chunk_offset + 8) >> 3;
-
- const float scale_b_0 = dequant_fp8_e4m3(sfb_shared[block_0]);
- const float scale_b_1 = dequant_fp8_e4m3(sfb_shared[block_1]);
-
- float b_vals[CHUNK_SIZE * 2];
-
- #pragma unroll
- for (int i = 0; i < 8; ++i) {
- const unsigned char b_val = b_shared[k_byte_base + chunk_offset + i];
- b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b_0;
- b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b_0;
- }
- #pragma unroll
- for (int i = 8; i < CHUNK_SIZE; ++i) {
- const unsigned char b_val = b_shared[k_byte_base + chunk_offset + i];
- b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b_1;
- b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b_1;
- }
-
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- if (!row_valid[r]) continue;
-
- const int m = base_m + r;
- const int64_t a_offset = m * a_s0 + (k_byte_base + chunk_offset) * a_s1 + l * a_s2;
-
- uint4 a_vec = __ldg(reinterpret_cast<const uint4*>(&a[a_offset]));
- unsigned char a_bytes[CHUNK_SIZE];
- *reinterpret_cast<uint4*>(&a_bytes[0]) = a_vec;
-
- const int kk_0 = block_0 / 4, kk4_0 = block_0 % 4;
- const int kk_1 = block_1 / 4, kk4_1 = block_1 % 4;
-
- const float scale_a_0 = dequant_fp8_e4m3(__ldg(&sfa[sfa_m_parts[r] + kk4_0 * sfa_s3 + kk_0 * sfa_s4]));
- const float scale_a_1 = dequant_fp8_e4m3(__ldg(&sfa[sfa_m_parts[r] + kk4_1 * sfa_s3 + kk_1 * sfa_s4]));
-
- #pragma unroll
- for (int i = 0; i < 8; ++i) {
- const unsigned char a_val = a_bytes[i];
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a_0, b_vals[i * 2], thread_acc[r]);
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a_0, b_vals[i * 2 + 1], thread_acc[r]);
- }
- #pragma unroll
- for (int i = 8; i < CHUNK_SIZE; ++i) {
- const unsigned char a_val = a_bytes[i];
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a_1, b_vals[i * 2], thread_acc[r]);
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a_1, b_vals[i * 2 + 1], thread_acc[r]);
- }
- }
- }
- }
-
- for (int k_byte = remainder_start + lane_id; k_byte < K_bytes; k_byte += warpSize) {
- const int k_block = k_byte >> 3;
- const unsigned char b_val = b_shared[k_byte];
- const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
- const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0) * scale_b;
- const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1) * scale_b;
+ // 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
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- const int m = base_m + r;
- if (m >= M) continue;
-
- const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;
- const unsigned char a_val = __ldg(&a[a_offset]);
-
- const int kk = k_block / 4, kk4 = k_block % 4;
- const int64_t sfa_offset = sfa_m_parts[r] + kk4 * sfa_s3 + kk * sfa_s4;
-
- if (sfa_offset >= 0) {
- const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a, b_fp4_0, thread_acc[r]);
- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a, b_fp4_1, thread_acc[r]);
- }
- }
+ 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;
+ 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;
+ 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* 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 sum0123 = __hadd2(sum01, sum23);
+ half2 sum4567 = __hadd2(sum45, sum67);
+
+ half2 sum_all = __hadd2(sum0123, sum4567);
+
+ float local_sum = __half2float(sum_all.x) + __half2float(sum_all.y);
+
+ acc += local_sum * scale;
}
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- float sum = thread_acc[r];
- #pragma unroll
- for (int mask = 16; mask > 0; mask >>= 1) {
- sum += __shfl_xor_sync(0xFFFFFFFF, sum, mask);
- }
-
- if (lane_id == 0) {
- const int m = base_m + r;
- if (m < M) {
- c[m * c_s0 + l * c_s2] = __float2half(sum);
- }
- }
+ acc = warp_reduce_sum(acc);
+ if (lane_id == 0) {
+ C[row * c_s0 + batch * c_s2] = __float2half(acc);
}
}
- torch::Tensor nvfp4_gemv(
- torch::Tensor a,
- torch::Tensor b,
- torch::Tensor sfa_permuted,
- torch::Tensor sfb_permuted,
+ torch::Tensor gemv_cuda(
+ torch::Tensor a, torch::Tensor b,
+ torch::Tensor sfa, torch::Tensor sfb,
torch::Tensor c
) {
- TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");
- TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
- TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
+ int M = a.size(0), K = a.size(1) * 2, L = a.size(2);
- static bool fp8_lut_initialized = false;
- if (!fp8_lut_initialized) {
- float host_fp8_lut[256];
- for (int i = 0; i < 256; ++i) {
- unsigned char fp8_bits = static_cast<unsigned char>(i);
- int sign = (fp8_bits >> 7) & 0x1;
- int exp = (fp8_bits >> 3) & 0xF;
- int mant = fp8_bits & 0x7;
-
- float val;
- if (exp == 0) {
- val = ldexpf(mant / 8.0f, -6);
- } else if (exp == 15) {
- val = 448.0f;
- } else {
- val = ldexpf(1.0f + mant / 8.0f, exp - 7);
- }
- host_fp8_lut[i] = sign ? -val : val;
- }
- cudaMemcpyToSymbol(fp8_e4m3_lut, host_fp8_lut, 256 * sizeof(float));
- fp8_lut_initialized = true;
- }
+ const int WARPS_PER_BLOCK = 8;
+ dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
+ dim3 block(32 * WARPS_PER_BLOCK);
- unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());
- unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());
- unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());
- unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());
-
- int M = a.size(0);
- int K_bytes = a.size(1);
- int L = a.size(2);
- int K = K_bytes * 2;
- int B_rows = b.size(0);
-
- int sfa_dim2 = sfa_permuted.size(2);
- int sfa_dim4 = sfa_permuted.size(4);
- int sfb_dim2 = sfb_permuted.size(2);
- int sfb_dim4 = sfb_permuted.size(4);
-
- dim3 block(32, 16);
- size_t smem_size = K_bytes + 4 + K / 16 + 1;
-
- // Adaptive configuration: favor parallelism for high-L cases
- int total_work = M * L;
- bool high_batch = (L >= 4);
-
- if (high_batch) {
- // Use ROWS_PER_WARP=2, VEC_SIZE=32 for better parallelism
- constexpr int ROWS_PER_WARP = 2;
- constexpr int VEC_SIZE = 32;
- const int rows_per_block = block.y * ROWS_PER_WARP;
- dim3 grid((M + rows_per_block - 1) / rows_per_block, L);
-
- nvfp4_gemv_kernel<ROWS_PER_WARP, VEC_SIZE><<<grid, block, smem_size>>>(
- a_ptr, b_ptr, sfa_ptr, sfb_ptr,
- reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
- M, K, L, B_rows,
- sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
- a.stride(0), a.stride(1), a.stride(2),
- b.stride(0), b.stride(1), b.stride(2),
- sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
- sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
- sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
- sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
- c.stride(0), c.stride(1), c.stride(2)
- );
- } else {
- // Use ROWS_PER_WARP=4, VEC_SIZE=32 for better arithmetic intensity
- constexpr int ROWS_PER_WARP = 4;
- constexpr int VEC_SIZE = 32;
- const int rows_per_block = block.y * ROWS_PER_WARP;
- dim3 grid((M + rows_per_block - 1) / rows_per_block, L);
-
- nvfp4_gemv_kernel<ROWS_PER_WARP, VEC_SIZE><<<grid, block, smem_size>>>(
- a_ptr, b_ptr, sfa_ptr, sfb_ptr,
- reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
- M, K, L, B_rows,
- sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
- a.stride(0), a.stride(1), a.stride(2),
- b.stride(0), b.stride(1), b.stride(2),
- sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
- sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
- sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
- sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
- c.stride(0), c.stride(1), c.stride(2)
- );
- }
-
- cudaError_t err = cudaGetLastError();
- TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));
-
+ gemv_v5<<<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;
}
"""
- nvfp4_gemv_cpp = """
- #include <torch/extension.h>
- torch::Tensor nvfp4_gemv(
- 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);
"""
- nvfp4_module = load_inline(
- name='nvfp4_gemv_p0_optimized',
- cpp_sources=nvfp4_gemv_cpp,
- cuda_sources=nvfp4_gemv_cuda,
- functions=['nvfp4_gemv'],
- verbose=False,
- extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
+ module = load_inline(
+ name='gemv_v5',
+ 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):
- """
- Adaptive P0 Optimized NVFP4 GEMV:
- - Adaptive ROWS_PER_WARP: 2 for L>=4 (parallelism), 4 for L<4 (intensity)
- - VEC_SIZE=32 with chunked processing to reduce register pressure
- - __ldg() for cache-optimized loads
- - Bank conflict padding
-
- Expected: 15-25% improvement across all cases
- """
- a, b, _, _, sfa_permuted, sfb_permuted, c = data
- return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)
No newline at end of file
+ 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)
No newline at end of file
scrolls · 501 diff lines total

Best evidence level for this revision: reported

JSON