Skip to content
KernelIndex
Search⌘K

submission 70699

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70699?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
88.0µs
#366 of 678
2025-11-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2495184b64e7ed10a0c1e1bbddfdeeb9f6d5a2e02f2086a9124a4871b5758d7a
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15

Techniques

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

fp4Adaptive P0 Optimized NVFP4 GEMV:
shared-memoryextern __shared__ unsigned char smem[];
vector-width = uint4__ldg(reinterpret_cast<const uint4*>(&b[b_offset_base]));

Kernel source

submission.py338 lines
import torch
from torch.utils.cpp_extension import load_inline

nvfp4_gemv_cuda = """
#include <cuda_fp16.h>
#include <cuda_runtime.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
};

__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 dequant_fp8_e4m3(unsigned char fp8_bits) {
    return fp8_e4m3_lut[fp8_bits];
}

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,
    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
) {
    extern __shared__ unsigned char smem[];
    unsigned char* b_shared = smem;
    unsigned char* sfb_shared = smem + (K / 2) + 4;
    
    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;
    
    const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;
    
    const int K_bytes = K / 2;
    const int K_blocks = K / 16;
    const int b_m = 0;
    
    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 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 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 in chunks to reduce register pressure
        constexpr int CHUNK_SIZE = VEC_SIZE / 2;
        
        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;
        
        #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]);
            }
        }
    }
    
    #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);
            }
        }
    }
}

torch::Tensor nvfp4_gemv(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa_permuted,
    torch::Tensor sfb_permuted,
    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");
    
    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;
    }
    
    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));
    
    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
);
"""

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']
)

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)
scrolls · 338 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 70452.

import torch
from torch.utils.cpp_extension import load_inline
- from task import input_t, output_t
nvfp4_gemv_cuda = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
- // FP4 E2M1 lookup table in constant memory for efficient broadcast to all threads
__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
+ 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
};
- // FP8 E4M3 lookup table in constant memory (256 entries = 1KB)
- // Precomputed on host for all possible FP8 E4M3 values
__constant__ float fp8_e4m3_lut[256];
__device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
- unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : (packed_val >> 4);
- return fp4_lut[fp4_bits];
+ return fp4_lut[(idx == 0) ? (packed_val & 0x0F) : (packed_val >> 4)];
}
__device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
return fp8_e4m3_lut[fp8_bits];
}
+ template<int ROWS_PER_WARP, int VEC_SIZE>
__global__ void nvfp4_gemv_kernel(
const unsigned char* __restrict__ a,
const unsigned char* __restrict__ b,
⋯ 10 unchanged lines
) {
extern __shared__ unsigned char smem[];
unsigned char* b_shared = smem;
- unsigned char* sfb_shared = smem + (K / 2);
+ unsigned char* sfb_shared = smem + (K / 2) + 4;
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;
- // Each warp now computes 2 output rows to balance arithmetic intensity vs dependencies
- constexpr int ROWS_PER_WARP = 2;
const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;
const int K_bytes = K / 2;
const int K_blocks = K / 16;
const int b_m = 0;
- // Vectorized loading of b into shared memory (16 bytes per thread)
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) {
+ 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;
- // Load 16 bytes as 2x uint64_t (assumes b_s1 stride allows this)
- if (b_s1 == 1) {
- *reinterpret_cast<uint64_t*>(&b_shared[byte_idx]) =
- *reinterpret_cast<const uint64_t*>(&b[b_offset_base]);
- *reinterpret_cast<uint64_t*>(&b_shared[byte_idx + 8]) =
- *reinterpret_cast<const uint64_t*>(&b[b_offset_base + 8]);
- } else {
- for (int i = 0; i < B_VEC_SIZE; ++i) {
- b_shared[byte_idx + i] = __ldg(&b[(byte_idx + i) * 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]);
⋯ 1 unchanged lines
}
}
- // Vectorized loading of sfb into shared memory
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;
⋯ 3 unchanged lines
}
__syncthreads();
- // Multiple accumulators - one per output row
float thread_acc[ROWS_PER_WARP];
#pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- thread_acc[r] = 0.0f;
- }
+ for (int r = 0; r < ROWS_PER_WARP; ++r) thread_acc[r] = 0.0f;
- // Precompute scale offset components for each row
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;
- if (m < M) {
+ 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;
⋯ 1 unchanged lines
}
}
- // Optimized vector size: process 16 consecutive bytes per thread per iteration
- // 16 bytes = 32 FP4 values = 2 NVFP4 scale blocks
- // This reduces iterations by 2x, amortizing loop overhead
- constexpr int VEC_SIZE = 16;
const int total_vec_iters = K_bytes / (warpSize * VEC_SIZE);
const int remainder_start = total_vec_iters * warpSize * VEC_SIZE;
- // Main vectorized loop - process VEC_SIZE bytes per thread per iteration
- // B vector is loaded ONCE and reused for ALL rows (key optimization!)
for (int iter = 0; iter < total_vec_iters; ++iter) {
const int k_byte_base = iter * warpSize * VEC_SIZE + lane_id * VEC_SIZE;
- // With VEC_SIZE=16, we span exactly 2 scale blocks
- const int first_block = k_byte_base >> 3;
- const int second_block = (k_byte_base + 8) >> 3;
+ // Process in chunks to reduce register pressure
+ constexpr int CHUNK_SIZE = VEC_SIZE / 2;
- // Load scales for both blocks
- const float scale_b_0 = dequant_fp8_e4m3(sfb_shared[first_block]);
- const float scale_b_1 = dequant_fp8_e4m3(sfb_shared[second_block]);
-
- // Dequantize b values for both blocks (shared across all rows)
- float b_vals[VEC_SIZE * 2]; // 2 FP4 values per byte
- #pragma unroll
- for (int i = 0; i < 8; ++i) {
- const unsigned char b_val = b_shared[k_byte_base + 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 < VEC_SIZE; ++i) {
- const unsigned char b_val = b_shared[k_byte_base + 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;
- }
-
- // Process each row with the preloaded b values
- #pragma unroll
- for (int r = 0; r < ROWS_PER_WARP; ++r) {
- const int m = base_m + r;
- if (m >= M) continue;
+ 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 int64_t a_offset = m * a_s0 + k_byte_base * a_s1 + l * a_s2;
+ const float scale_b_0 = dequant_fp8_e4m3(sfb_shared[block_0]);
+ const float scale_b_1 = dequant_fp8_e4m3(sfb_shared[block_1]);
- // Vectorized load for this row's a values (16 bytes via 2x uint64_t)
- uint64_t a_vec_0 = *reinterpret_cast<const uint64_t*>(&a[a_offset]);
- uint64_t a_vec_1 = *reinterpret_cast<const uint64_t*>(&a[a_offset + 8]);
- unsigned char a_bytes[16];
- *reinterpret_cast<uint64_t*>(&a_bytes[0]) = a_vec_0;
- *reinterpret_cast<uint64_t*>(&a_bytes[8]) = a_vec_1;
+ float b_vals[CHUNK_SIZE * 2];
- // Load scales for this row's a values (2 blocks)
- const int kk_0 = first_block / 4;
- const int kk4_0 = first_block % 4;
- const int64_t sfa_offset_0 = sfa_m_parts[r] + kk4_0 * sfa_s3 + kk_0 * sfa_s4;
- const float scale_a_0 = (sfa_offset_0 >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset_0])) : 0.0f;
-
- const int kk_1 = second_block / 4;
- const int kk4_1 = second_block % 4;
- const int64_t sfa_offset_1 = sfa_m_parts[r] + kk4_1 * sfa_s3 + kk_1 * sfa_s4;
- const float scale_a_1 = (sfa_offset_1 >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset_1])) : 0.0f;
-
- // Compute dot products using preloaded b values
- // First block (bytes 0-7)
#pragma unroll
for (int i = 0; i < 8; ++i) {
- const unsigned char a_val = a_bytes[i];
- const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a_0;
- const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a_0;
-
- thread_acc[r] = __fmaf_rn(a_fp4_0, b_vals[i * 2], thread_acc[r]);
- thread_acc[r] = __fmaf_rn(a_fp4_1, b_vals[i * 2 + 1], thread_acc[r]);
+ 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;
+ }
- // Second block (bytes 8-15)
#pragma unroll
- for (int i = 8; i < VEC_SIZE; ++i) {
- const unsigned char a_val = a_bytes[i];
- const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a_1;
- const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a_1;
+ for (int r = 0; r < ROWS_PER_WARP; ++r) {
+ if (!row_valid[r]) continue;
- thread_acc[r] = __fmaf_rn(a_fp4_0, b_vals[i * 2], thread_acc[r]);
- thread_acc[r] = __fmaf_rn(a_fp4_1, b_vals[i * 2 + 1], thread_acc[r]);
+ 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]);
+ }
}
}
}
- // Handle remainder bytes
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;
⋯ 7 unchanged lines
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;
- const int kk4 = k_block % 4;
+ 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]));
- const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a;
- const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a;
-
- thread_acc[r] = __fmaf_rn(a_fp4_0, b_fp4_0, thread_acc[r]);
- thread_acc[r] = __fmaf_rn(a_fp4_1, b_fp4_1, thread_acc[r]);
+ 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]);
}
}
}
- // Warp reduction for each output row separately
#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);
⋯ 2 unchanged lines
if (lane_id == 0) {
const int m = base_m + r;
if (m < M) {
- const int64_t c_offset = m * c_s0 + l * c_s2;
- c[c_offset] = __float2half(sum);
+ c[m * c_s0 + l * c_s2] = __float2half(sum);
}
}
}
⋯ 10 unchanged lines
TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
- // Initialize FP8 E4M3 lookup table once
static bool fp8_lut_initialized = false;
if (!fp8_lut_initialized) {
float host_fp8_lut[256];
⋯ 13 unchanged lines
}
host_fp8_lut[i] = sign ? -val : val;
}
-
- // Copy to constant memory
cudaMemcpyToSymbol(fp8_e4m3_lut, host_fp8_lut, 256 * sizeof(float));
fp8_lut_initialized = true;
}
⋯ 15 unchanged lines
int sfb_dim4 = sfb_permuted.size(4);
dim3 block(32, 16);
- // Each warp now computes 2 rows, so adjust grid accordingly
- constexpr int ROWS_PER_WARP = 2;
- const int rows_per_block = block.y * ROWS_PER_WARP; // 16 warps * 2 rows = 32 rows per block
- dim3 grid((M + rows_per_block - 1) / rows_per_block, L);
- size_t smem_size = K_bytes + K / 16;
+ size_t smem_size = K_bytes + 4 + K / 16 + 1;
- nvfp4_gemv_kernel<<<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)
- );
+ // 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));
⋯ 13 unchanged lines
"""
nvfp4_module = load_inline(
- name='nvfp4_gemv_vec16',
+ name='nvfp4_gemv_p0_optimized',
cpp_sources=nvfp4_gemv_cpp,
cuda_sources=nvfp4_gemv_cuda,
functions=['nvfp4_gemv'],
⋯ 1 unchanged lines
extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
)
- def custom_kernel(data: input_t) -> output_t:
+ def custom_kernel(data):
"""
- NVFP4 GEMV optimized for CUDA Cores (not Tensor Cores)
+ 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
- Why NOT using Tensor Cores:
- - GEMV is M×K @ K×1 (matrix × vector)
- - Tensor Cores require matrix×matrix (e.g., 16×16×16 tiles)
- - Vector dimension (N=1) doesn't map well to TC tile sizes
- - CUDA Core approach is more efficient for true GEMV operations
-
- Current optimizations:
- 1. Vectorized loads: 4-byte chunks via uint32_t
- 2. Constant memory LUTs: FP4 (16 entries) + FP8 (256 entries)
- 3. Hoisted scale loads: Load once per iteration (75% reduction)
- 4. Explicit FMA: __fmaf_rn for maximum FMA unit utilization
- 5. Zero branch divergence: All lookups via constant cache
-
- Performance progression on B200 Blackwell:
- - Baseline: 1.0x
- - + FP4 LUT: 2.0x
- - + FP8 LUT + Hoisted scales: ~4.0x (estimated)
- - + FMA instructions: 4.4-4.6x (expected)
-
- Profiling shows: ALU 50-55%, FMA 21-24%, TC 0% (intentional)
- Primary bottleneck: MIO Throttle (memory-bound, as expected for GEMV)
+ 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
scrolls · 432 diff lines total

Best evidence level for this revision: reported

JSON