Skip to content
KernelIndex
Search⌘K

submission 75888

mdouglas · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_cuda.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-75888?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
72.3µs
#332 of 678
2025-11-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9db8cb51bc93b7e5f7feae47e94a8c0a95246da6be54c8197a1ecab3099b3553
license declaredunknown
license concludedunknown
authorsmdouglas
imported2026-08-15

Techniques

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

fp4const uint8_t* __restrict__ a, // [M, K//2] packed FP4 (2 per byte)
fp8const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)
shared-memoryextern __shared__ uint8_t smem[];
vector-width = uint4reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];

Kernel source

submission_cuda.py416 lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# CUDA kernel code for NVFP4 block-scaled GEMV
# Uses native Blackwell (sm_100a) hardware intrinsics for FP4/FP8 conversion
cuda_source = """
#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>

// NVFP4 is 4-bit float (e2m1): 1 sign bit, 2 exponent bits, 1 mantissa bit
// Stored as 2 values per byte
// Scale factors are FP8 (e4m3) for every 16 FP4 values

// Warp reduction using shuffle - optimized for Blackwell
__device__ __forceinline__ float warp_reduce_sum(float val) {
    #pragma unroll
    for (int offset = 16; offset > 0; offset /= 2) {
        val += __shfl_down_sync(0xffffffff, val, offset);
    }
    return val;
}

// Batched kernel: processes all L batches in one launch for better efficiency.
// Each block handles one M value, each warp within handles one L batch.
// Inputs have native PyTorch strides from .permute() - K dimension has stride 1.
__global__ void nvfp4_gemv_batched_kernel(
    const uint8_t* __restrict__ a,      // [M, K//2, L] with strides (K_half, 1, M*K_half)
    const uint8_t* __restrict__ b,      // [N, K//2, L] with strides (K_half, 1, N*K_half), N=128 padded
    const __nv_fp8_e4m3* __restrict__ sfa,    // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)
    const __nv_fp8_e4m3* __restrict__ sfb,    // [N, K//16, L] with strides (K_div_16, 1, N*K_div_16)
    half* __restrict__ c,                // [M, 1, L] output FP16
    int M,
    int K,
    int L
) {
    // Shared memory layout: B vectors [L, K/2], sfb [L, K/16]
    extern __shared__ uint8_t smem[];
    int K_half = K / 2;
    int K_div_16 = K / 16;

    uint8_t* sb = smem;  // B vectors: L × K/2 bytes
    __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + L * K_half);  // Scale factors B: L × K/16

    int tid = threadIdx.x;
    int warp_id = threadIdx.x / 32;
    int lane = threadIdx.x % 32;
    int m = blockIdx.x;  // Each block handles one M value

    if (m >= M) return;

    const int N_padded = 128;  // B is padded to 128 rows for torch._scaled_mm

    // Cooperatively load all L B vectors into shared memory
    // B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half (strides: K_half, 1, N_padded*K_half)
    // We only need n=0 (the actual vector, rest is padding)
    for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
        int k = kl / L;
        int l = kl % L;
        // b[0, k, l] at offset: 0*K_half + k + l*N_padded*K_half
        sb[l * K_half + k] = b[k + l * N_padded * K_half];
    }

    // Cooperatively load scale factors for B
    for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
        int k = kl / L;
        int l = kl % L;
        // sfb[0, k, l] at offset: 0*K_div_16 + k + l*N_padded*K_div_16
        ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
    }

    __syncthreads();

    // Each block processes multiple M rows, each warp processes multiple (m, l) pairs
    const int M_ROWS_PER_BLOCK = 4;
    const int WARPS_PER_BLOCK = blockDim.x / 32;
    int m_start = blockIdx.x * M_ROWS_PER_BLOCK;

    // Each warp processes all L batches for M_ROWS_PER_BLOCK / WARPS_PER_BLOCK M rows
    int m_rows_per_warp = (M_ROWS_PER_BLOCK + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK;

    for (int m_idx = 0; m_idx < m_rows_per_warp; m_idx++) {
        int m = m_start + warp_id + m_idx * WARPS_PER_BLOCK;
        if (m >= M) break;

        // Process all L batches for this M row
        for (int l = 0; l < L; l++) {
            float sum = 0.0f;

            // K dimension has stride 1, enabling coalesced access.
            const int a_base = m * K_half + l * M * K_half;
            const int sfa_base = m * K_div_16 + l * M * K_div_16;
            const uint8_t* sb_row = &sb[l * K_half];
            const __nv_fp8_e4m3* ssfb_row = &ssfb[l * K_div_16];

            // Process in blocks of 8 bytes - each scale factor covers 8 bytes (16 FP4 values).
            for (int scale_block = lane; scale_block < K_div_16; scale_block += 32) {
                half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block].__x, __NV_E4M3).x);
                half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block].__x, __NV_E4M3).x);
                half combined_scale = scale_a * scale_b;
                __half2 scale2 = __half2half2(combined_scale);

                // Load all 8 bytes at once using uint2
                int k_byte_base = scale_block * 8;
                const uint2 a_data = *reinterpret_cast<const uint2*>(&a[a_base + k_byte_base]);
                const uint2 b_data = *reinterpret_cast<const uint2*>(&sb_row[k_byte_base]);

                const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
                const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);

                // Process all 8 bytes with the same scale
                #pragma unroll
                for (int i = 0; i < 8; i++) {
                    __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
                    __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);

                    __half2 products = __hmul2(a_vals, b_vals);
                    __half2 scaled = __hmul2(products, scale2);

                    sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
                    sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
                }
            }

            sum = warp_reduce_sum(sum);

            // c_ref has shape [M, 1, L] with strides (1, 1, M) from permute
            // So c_ref[m, 0, l] is at linear offset: m + l*M
            if (lane == 0) {
                c[m + l * M] = __float2half(sum);
            }
        }
    }
}

__global__ void nvfp4_gemv_kernel(
    const uint8_t* __restrict__ a,      // [M, K//2] packed FP4 (2 per byte)
    const uint8_t* __restrict__ b,      // [1, K//2] packed FP4
    const __nv_fp8_e4m3* __restrict__ sfa,    // [M, K//16] FP8 scale factors for A
    const __nv_fp8_e4m3* __restrict__ sfb,    // [1, K//16] FP8 scale factors for B
    half* __restrict__ c,                // [M, 1] output FP16
    int M,
    int K
) {
    // 8 warps per M row for better memory latency hiding
    // 2 M rows per block
    const int WARPS_PER_M_ROW = 8;
    const int WARPS_PER_BLOCK = blockDim.x / 32;
    const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;  // = 2

    // Shared memory: B vector, scale factors, and partial sums for reduction
    extern __shared__ uint8_t smem[];
    uint8_t* sb = smem;
    __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
    float* partial_sums = reinterpret_cast<float*>(ssfb + K/16);

    int tid = threadIdx.x;
    int warp_id = tid / 32;
    int lane = tid % 32;

    int K_half = K / 2;
    int K_div_16 = K / 16;

    // Cooperatively load B vector into shared memory (vectorized)
    int num_vec_loads = K_half / 16;
    for (int i = tid; i < num_vec_loads; i += blockDim.x) {
        reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];
    }
    int vec_bytes = num_vec_loads * 16;
    for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
        sb[i] = b[i];
    }

    // Cooperatively load scale factors for B
    for (int i = tid; i < K_div_16; i += blockDim.x) {
        ssfb[i] = sfb[i];
    }

    __syncthreads();

    // Each block processes M_ROWS_PER_BLOCK M rows
    int m_base = blockIdx.x * M_ROWS_PER_BLOCK;

    // Which M row within the block does this warp contribute to?
    int m_local = warp_id / WARPS_PER_M_ROW;
    int m = m_base + m_local;

    if (m >= M) return;

    // Which K chunk does this warp handle?
    int warp_in_m_group = warp_id % WARPS_PER_M_ROW;

    // Process by scale blocks instead of bytes - each scale covers 8 bytes (16 FP4 values)
    int scales_per_warp = K_div_16 / WARPS_PER_M_ROW;  // 1024 / 8 = 128 scales per warp
    int scale_start = warp_in_m_group * scales_per_warp;
    int scale_end = scale_start + scales_per_warp;

    float sum = 0.0f;

    // Loop over scale blocks - each iteration processes 8 bytes covered by one scale
    // With 8 warps: 128 scales / 32 threads = 4 iterations per thread (down from 32!)
    for (int scale_block = scale_start + lane; scale_block < scale_end; scale_block += 32) {
        // Load scale factors ONCE for this block of 8 bytes
        half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block].__x, __NV_E4M3).x);
        half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block].__x, __NV_E4M3).x);
        half combined_scale = scale_a * scale_b;
        __half2 scale2 = __half2half2(combined_scale);

        // Process all 8 bytes (8 fp4x2 pairs) covered by this scale factor
        int k_byte_base = scale_block * 8;

        // Load 8 bytes at once using uint2, then reinterpret as fp4x2 array
        const uint2 a_data = *reinterpret_cast<const uint2*>(&a[m * K_half + k_byte_base]);
        const uint2 b_data = *reinterpret_cast<const uint2*>(&sb[k_byte_base]);

        const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
        const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);

        // Process all 8 bytes with the same scale
        #pragma unroll
        for (int i = 0; i < 8; i++) {
            __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
            __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);

            __half2 products = __hmul2(a_vals, b_vals);
            __half2 scaled = __hmul2(products, scale2);

            sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
            sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
        }
    }

    // Intra-warp reduction
    sum = warp_reduce_sum(sum);

    // Store partial sum to shared memory
    if (lane == 0) {
        partial_sums[warp_id] = sum;
    }

    __syncthreads();

    // Final reduction: first warp of each M group reduces the partial sums
    if (warp_in_m_group == 0 && lane < WARPS_PER_M_ROW) {
        float final_sum = partial_sums[m_local * WARPS_PER_M_ROW + lane];
        // Reduce across the 8 partial sums
        #pragma unroll
        for (int offset = 4; offset > 0; offset /= 2) {
            final_sum += __shfl_down_sync(0xffffffff, final_sum, offset);
        }

        if (lane == 0) {
            c[m] = __float2half(final_sum);
        }
    }
}

void nvfp4_gemv_cuda(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M,
    int K
) {
    // 8 warps per M row, 2 M rows per block
    const int WARPS_PER_M_ROW = 8;
    const int WARPS_PER_BLOCK = 16;
    const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;  // = 2
    const int threads = WARPS_PER_BLOCK * 32;  // 512 threads
    const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

    // Shared memory: B vector + sfb + partial sums
    const int smem_size = K / 2 + K / 16 + WARPS_PER_BLOCK * sizeof(float);

    // Get current CUDA stream from PyTorch
    cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());

    nvfp4_gemv_kernel<<<blocks, threads, smem_size, stream>>>(
        a.data_ptr<uint8_t>(),
        b.data_ptr<uint8_t>(),
        reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
        reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
        reinterpret_cast<half*>(c.data_ptr<at::Half>()),
        M, K
    );

    // Check for kernel launch errors
    AT_CUDA_CHECK(cudaGetLastError());
}

void nvfp4_gemv_batched_cuda(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M,
    int K,
    int L
) {
    // Each block handles 4 M rows, 8 warps process all (m, l) pairs
    const int M_ROWS_PER_BLOCK = 4;
    const int threads = 256;
    const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

    // Shared memory: B vectors (L × K/2) and sfb (L × K/16)
    const int K_half = K / 2;
    const int K_div_16 = K / 16;
    const int smem_size = L * K_half + L * K_div_16;

    // Get current CUDA stream from PyTorch
    cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());

    nvfp4_gemv_batched_kernel<<<blocks, threads, smem_size, stream>>>(
        a.data_ptr<uint8_t>(),
        b.data_ptr<uint8_t>(),
        reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
        reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
        reinterpret_cast<half*>(c.data_ptr<at::Half>()),
        M, K, L
    );

    // Check for kernel launch errors
    AT_CUDA_CHECK(cudaGetLastError());
}
"""

cpp_source = """
void nvfp4_gemv_cuda(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M,
    int K
);

void nvfp4_gemv_batched_cuda(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M,
    int K,
    int L
);
"""

# Compile the CUDA extension inline
nvfp4_gemv_module = load_inline(
    name='nvfp4_gemv',
    cpp_sources=[cpp_source],
    cuda_sources=[cuda_source],
    functions=['nvfp4_gemv_cuda', 'nvfp4_gemv_batched_cuda'],
    verbose=True,
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '-arch=sm_100a',
        '--std=c++17',
        '-U__CUDA_NO_HALF_OPERATORS__',  # Enable half operators
        '-U__CUDA_NO_HALF_CONVERSIONS__',  # Enable half conversions
    ],
)

def custom_kernel(data: input_t) -> output_t:
    """
    Custom CUDA implementation of NVFP4 block-scaled GEMV.
    Uses separate kernels optimized for L=1 and L>1 cases.
    """
    import torch
    a_ref, b_ref, sfa, sfb, _, _, c_ref = data

    M, K_half, L = a_ref.shape
    K = K_half * 2  # Each byte contains 2 FP4 values

    if L == 1:
        # Use single-batch optimized kernel.
        # Pass c_ref directly - kernel writes c[m] which maps to c_ref[m, 0, 0].
        a_bytes = a_ref[:, :, 0].view(torch.uint8).contiguous()
        b_bytes = b_ref[0, :, 0].view(torch.uint8).contiguous()

        nvfp4_gemv_module.nvfp4_gemv_cuda(
            a_bytes,
            b_bytes,
            sfa[:, :, 0].contiguous(),
            sfb[0, :, 0].contiguous(),
            c_ref,
            M, K
        )
    else:
        # Use batched kernel for L>1 - processes all batches in one launch.
        # Original layout already has K contiguous (stride 1) from creation
        # as (L, M, K).permute(1, 2, 0), so no additional permute is needed.
        a_bytes = a_ref.view(torch.uint8)
        b_bytes = b_ref.view(torch.uint8)

        nvfp4_gemv_module.nvfp4_gemv_batched_cuda(
            a_bytes,
            b_bytes,
            sfa,
            sfb,
            c_ref,
            M, K, L
        )

    return c_ref
scrolls · 416 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 71049.

- import torch
+ from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
- torch._dynamo.config.cache_size_limit = 32
+ # CUDA kernel code for NVFP4 block-scaled GEMV
+ # Uses native Blackwell (sm_100a) hardware intrinsics for FP4/FP8 conversion
+ cuda_source = """
+ #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>
- # Convert all scale factors to blocked formats
+ // NVFP4 is 4-bit float (e2m1): 1 sign bit, 2 exponent bits, 1 mantissa bit
+ // Stored as 2 values per byte
+ // Scale factors are FP8 (e4m3) for every 16 FP4 values
- @torch.compile(dynamic=False, fullgraph=True)
- def to_blocked_3d(input_matrix):
- # input_matrix is rows x cols x l
- rows, cols, l = input_matrix.shape
+ // Warp reduction using shuffle - optimized for Blackwell
+ __device__ __forceinline__ float warp_reduce_sum(float val) {
+ #pragma unroll
+ for (int offset = 16; offset > 0; offset /= 2) {
+ val += __shfl_down_sync(0xffffffff, val, offset);
+ }
+ return val;
+ }
- data = input_matrix.permute(2, 0, 1)
+ // Batched kernel: processes all L batches in one launch for better efficiency.
+ // Each block handles one M value, each warp within handles one L batch.
+ // Inputs have native PyTorch strides from .permute() - K dimension has stride 1.
+ __global__ void nvfp4_gemv_batched_kernel(
+ const uint8_t* __restrict__ a, // [M, K//2, L] with strides (K_half, 1, M*K_half)
+ const uint8_t* __restrict__ b, // [N, K//2, L] with strides (K_half, 1, N*K_half), N=128 padded
+ const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)
+ const __nv_fp8_e4m3* __restrict__ sfb, // [N, K//16, L] with strides (K_div_16, 1, N*K_div_16)
+ half* __restrict__ c, // [M, 1, L] output FP16
+ int M,
+ int K,
+ int L
+ ) {
+ // Shared memory layout: B vectors [L, K/2], sfb [L, K/16]
+ extern __shared__ uint8_t smem[];
+ int K_half = K / 2;
+ int K_div_16 = K / 16;
- return data.view(l, rows // 128, 128, cols // 4, 4) \
- .transpose(2, 3) \
- .reshape(l, -1, 4, 32, 4) \
- .transpose(2, 3) \
- .flatten(1)
+ uint8_t* sb = smem; // B vectors: L × K/2 bytes
+ __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + L * K_half); // Scale factors B: L × K/16
- @torch.compile(dynamic=False, mode="max-autotune-no-cudagraphs")
- def batched_gemv_impl(a_ref, b_ref, c_ref, sfa, sfb):
- _, _, l = b_ref.shape
+ int tid = threadIdx.x;
+ int warp_id = threadIdx.x / 32;
+ int lane = threadIdx.x % 32;
+ int m = blockIdx.x; // Each block handles one M value
- sfa_blocked = to_blocked_3d(sfa)
- sfb_blocked = to_blocked_3d(sfb)
+ if (m >= M) return;
- for l_idx in range(l):
- c_ref[:, 0, l_idx] = torch._scaled_mm(
- a_ref[..., l_idx],
- b_ref[..., l_idx].t(),
- sfa_blocked[l_idx, ...],
- sfb_blocked[l_idx, ...],
- bias=None,
- out_dtype=torch.float16,
- )[:, 0]
+ const int N_padded = 128; // B is padded to 128 rows for torch._scaled_mm
- return c_ref
+ // Cooperatively load all L B vectors into shared memory
+ // B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half (strides: K_half, 1, N_padded*K_half)
+ // We only need n=0 (the actual vector, rest is padding)
+ for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
+ int k = kl / L;
+ int l = kl % L;
+ // b[0, k, l] at offset: 0*K_half + k + l*N_padded*K_half
+ sb[l * K_half + k] = b[k + l * N_padded * K_half];
+ }
- def custom_kernel(
- data: input_t,
- ) -> output_t:
+ // Cooperatively load scale factors for B
+ for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
+ int k = kl / L;
+ int l = kl % L;
+ // sfb[0, k, l] at offset: 0*K_div_16 + k + l*N_padded*K_div_16
+ ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
+ }
+
+ __syncthreads();
+
+ // Each block processes multiple M rows, each warp processes multiple (m, l) pairs
+ const int M_ROWS_PER_BLOCK = 4;
+ const int WARPS_PER_BLOCK = blockDim.x / 32;
+ int m_start = blockIdx.x * M_ROWS_PER_BLOCK;
+
+ // Each warp processes all L batches for M_ROWS_PER_BLOCK / WARPS_PER_BLOCK M rows
+ int m_rows_per_warp = (M_ROWS_PER_BLOCK + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK;
+
+ for (int m_idx = 0; m_idx < m_rows_per_warp; m_idx++) {
+ int m = m_start + warp_id + m_idx * WARPS_PER_BLOCK;
+ if (m >= M) break;
+
+ // Process all L batches for this M row
+ for (int l = 0; l < L; l++) {
+ float sum = 0.0f;
+
+ // K dimension has stride 1, enabling coalesced access.
+ const int a_base = m * K_half + l * M * K_half;
+ const int sfa_base = m * K_div_16 + l * M * K_div_16;
+ const uint8_t* sb_row = &sb[l * K_half];
+ const __nv_fp8_e4m3* ssfb_row = &ssfb[l * K_div_16];
+
+ // Process in blocks of 8 bytes - each scale factor covers 8 bytes (16 FP4 values).
+ for (int scale_block = lane; scale_block < K_div_16; scale_block += 32) {
+ half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block].__x, __NV_E4M3).x);
+ half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block].__x, __NV_E4M3).x);
+ half combined_scale = scale_a * scale_b;
+ __half2 scale2 = __half2half2(combined_scale);
+
+ // Load all 8 bytes at once using uint2
+ int k_byte_base = scale_block * 8;
+ const uint2 a_data = *reinterpret_cast<const uint2*>(&a[a_base + k_byte_base]);
+ const uint2 b_data = *reinterpret_cast<const uint2*>(&sb_row[k_byte_base]);
+
+ const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
+ const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);
+
+ // Process all 8 bytes with the same scale
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
+ __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);
+
+ __half2 products = __hmul2(a_vals, b_vals);
+ __half2 scaled = __hmul2(products, scale2);
+
+ sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
+ sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
+ }
+ }
+
+ sum = warp_reduce_sum(sum);
+
+ // c_ref has shape [M, 1, L] with strides (1, 1, M) from permute
+ // So c_ref[m, 0, l] is at linear offset: m + l*M
+ if (lane == 0) {
+ c[m + l * M] = __float2half(sum);
+ }
+ }
+ }
+ }
+
+ __global__ void nvfp4_gemv_kernel(
+ const uint8_t* __restrict__ a, // [M, K//2] packed FP4 (2 per byte)
+ const uint8_t* __restrict__ b, // [1, K//2] packed FP4
+ const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16] FP8 scale factors for A
+ const __nv_fp8_e4m3* __restrict__ sfb, // [1, K//16] FP8 scale factors for B
+ half* __restrict__ c, // [M, 1] output FP16
+ int M,
+ int K
+ ) {
+ // 8 warps per M row for better memory latency hiding
+ // 2 M rows per block
+ const int WARPS_PER_M_ROW = 8;
+ const int WARPS_PER_BLOCK = blockDim.x / 32;
+ const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2
+
+ // Shared memory: B vector, scale factors, and partial sums for reduction
+ extern __shared__ uint8_t smem[];
+ uint8_t* sb = smem;
+ __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
+ float* partial_sums = reinterpret_cast<float*>(ssfb + K/16);
+
+ int tid = threadIdx.x;
+ int warp_id = tid / 32;
+ int lane = tid % 32;
+
+ int K_half = K / 2;
+ int K_div_16 = K / 16;
+
+ // Cooperatively load B vector into shared memory (vectorized)
+ int num_vec_loads = K_half / 16;
+ for (int i = tid; i < num_vec_loads; i += blockDim.x) {
+ reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];
+ }
+ int vec_bytes = num_vec_loads * 16;
+ for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
+ sb[i] = b[i];
+ }
+
+ // Cooperatively load scale factors for B
+ for (int i = tid; i < K_div_16; i += blockDim.x) {
+ ssfb[i] = sfb[i];
+ }
+
+ __syncthreads();
+
+ // Each block processes M_ROWS_PER_BLOCK M rows
+ int m_base = blockIdx.x * M_ROWS_PER_BLOCK;
+
+ // Which M row within the block does this warp contribute to?
+ int m_local = warp_id / WARPS_PER_M_ROW;
+ int m = m_base + m_local;
+
+ if (m >= M) return;
+
+ // Which K chunk does this warp handle?
+ int warp_in_m_group = warp_id % WARPS_PER_M_ROW;
+
+ // Process by scale blocks instead of bytes - each scale covers 8 bytes (16 FP4 values)
+ int scales_per_warp = K_div_16 / WARPS_PER_M_ROW; // 1024 / 8 = 128 scales per warp
+ int scale_start = warp_in_m_group * scales_per_warp;
+ int scale_end = scale_start + scales_per_warp;
+
+ float sum = 0.0f;
+
+ // Loop over scale blocks - each iteration processes 8 bytes covered by one scale
+ // With 8 warps: 128 scales / 32 threads = 4 iterations per thread (down from 32!)
+ for (int scale_block = scale_start + lane; scale_block < scale_end; scale_block += 32) {
+ // Load scale factors ONCE for this block of 8 bytes
+ half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block].__x, __NV_E4M3).x);
+ half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block].__x, __NV_E4M3).x);
+ half combined_scale = scale_a * scale_b;
+ __half2 scale2 = __half2half2(combined_scale);
+
+ // Process all 8 bytes (8 fp4x2 pairs) covered by this scale factor
+ int k_byte_base = scale_block * 8;
+
+ // Load 8 bytes at once using uint2, then reinterpret as fp4x2 array
+ const uint2 a_data = *reinterpret_cast<const uint2*>(&a[m * K_half + k_byte_base]);
+ const uint2 b_data = *reinterpret_cast<const uint2*>(&sb[k_byte_base]);
+
+ const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
+ const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);
+
+ // Process all 8 bytes with the same scale
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
+ __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);
+
+ __half2 products = __hmul2(a_vals, b_vals);
+ __half2 scaled = __hmul2(products, scale2);
+
+ sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
+ sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
+ }
+ }
+
+ // Intra-warp reduction
+ sum = warp_reduce_sum(sum);
+
+ // Store partial sum to shared memory
+ if (lane == 0) {
+ partial_sums[warp_id] = sum;
+ }
+
+ __syncthreads();
+
+ // Final reduction: first warp of each M group reduces the partial sums
+ if (warp_in_m_group == 0 && lane < WARPS_PER_M_ROW) {
+ float final_sum = partial_sums[m_local * WARPS_PER_M_ROW + lane];
+ // Reduce across the 8 partial sums
+ #pragma unroll
+ for (int offset = 4; offset > 0; offset /= 2) {
+ final_sum += __shfl_down_sync(0xffffffff, final_sum, offset);
+ }
+
+ if (lane == 0) {
+ c[m] = __float2half(final_sum);
+ }
+ }
+ }
+
+ void nvfp4_gemv_cuda(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa,
+ torch::Tensor sfb,
+ torch::Tensor c,
+ int M,
+ int K
+ ) {
+ // 8 warps per M row, 2 M rows per block
+ const int WARPS_PER_M_ROW = 8;
+ const int WARPS_PER_BLOCK = 16;
+ const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2
+ const int threads = WARPS_PER_BLOCK * 32; // 512 threads
+ const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
+
+ // Shared memory: B vector + sfb + partial sums
+ const int smem_size = K / 2 + K / 16 + WARPS_PER_BLOCK * sizeof(float);
+
+ // Get current CUDA stream from PyTorch
+ cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());
+
+ nvfp4_gemv_kernel<<<blocks, threads, smem_size, stream>>>(
+ a.data_ptr<uint8_t>(),
+ b.data_ptr<uint8_t>(),
+ reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
+ reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
+ reinterpret_cast<half*>(c.data_ptr<at::Half>()),
+ M, K
+ );
+
+ // Check for kernel launch errors
+ AT_CUDA_CHECK(cudaGetLastError());
+ }
+
+ void nvfp4_gemv_batched_cuda(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa,
+ torch::Tensor sfb,
+ torch::Tensor c,
+ int M,
+ int K,
+ int L
+ ) {
+ // Each block handles 4 M rows, 8 warps process all (m, l) pairs
+ const int M_ROWS_PER_BLOCK = 4;
+ const int threads = 256;
+ const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
+
+ // Shared memory: B vectors (L × K/2) and sfb (L × K/16)
+ const int K_half = K / 2;
+ const int K_div_16 = K / 16;
+ const int smem_size = L * K_half + L * K_div_16;
+
+ // Get current CUDA stream from PyTorch
+ cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());
+
+ nvfp4_gemv_batched_kernel<<<blocks, threads, smem_size, stream>>>(
+ a.data_ptr<uint8_t>(),
+ b.data_ptr<uint8_t>(),
+ reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
+ reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
+ reinterpret_cast<half*>(c.data_ptr<at::Half>()),
+ M, K, L
+ );
+
+ // Check for kernel launch errors
+ AT_CUDA_CHECK(cudaGetLastError());
+ }
+ """
+
+ cpp_source = """
+ void nvfp4_gemv_cuda(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa,
+ torch::Tensor sfb,
+ torch::Tensor c,
+ int M,
+ int K
+ );
+
+ void nvfp4_gemv_batched_cuda(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa,
+ torch::Tensor sfb,
+ torch::Tensor c,
+ int M,
+ int K,
+ int L
+ );
+ """
+
+ # Compile the CUDA extension inline
+ nvfp4_gemv_module = load_inline(
+ name='nvfp4_gemv',
+ cpp_sources=[cpp_source],
+ cuda_sources=[cuda_source],
+ functions=['nvfp4_gemv_cuda', 'nvfp4_gemv_batched_cuda'],
+ verbose=True,
+ extra_cuda_cflags=[
+ '-O3',
+ '--use_fast_math',
+ '-arch=sm_100a',
+ '--std=c++17',
+ '-U__CUDA_NO_HALF_OPERATORS__', # Enable half operators
+ '-U__CUDA_NO_HALF_CONVERSIONS__', # Enable half conversions
+ ],
+ )
+
+ def custom_kernel(data: input_t) -> output_t:
"""
- PyTorch reference implementation of NVFP4 block-scaled GEMV.
+ Custom CUDA implementation of NVFP4 block-scaled GEMV.
+ Uses separate kernels optimized for L=1 and L>1 cases.
"""
+ import torch
a_ref, b_ref, sfa, sfb, _, _, c_ref = data
- # a_ref is [m, k//2, l]
- # b_ref is [n, k//2, l], n=1 padded to n=128
- # c_ref is [m, 1, l]
- return batched_gemv_impl(
- a_ref,
- b_ref,
- c_ref,
- sfa,
- sfb,
- )
+ M, K_half, L = a_ref.shape
+ K = K_half * 2 # Each byte contains 2 FP4 values
+ if L == 1:
+ # Use single-batch optimized kernel.
+ # Pass c_ref directly - kernel writes c[m] which maps to c_ref[m, 0, 0].
+ a_bytes = a_ref[:, :, 0].view(torch.uint8).contiguous()
+ b_bytes = b_ref[0, :, 0].view(torch.uint8).contiguous()
+
+ nvfp4_gemv_module.nvfp4_gemv_cuda(
+ a_bytes,
+ b_bytes,
+ sfa[:, :, 0].contiguous(),
+ sfb[0, :, 0].contiguous(),
+ c_ref,
+ M, K
+ )
+ else:
+ # Use batched kernel for L>1 - processes all batches in one launch.
+ # Original layout already has K contiguous (stride 1) from creation
+ # as (L, M, K).permute(1, 2, 0), so no additional permute is needed.
+ a_bytes = a_ref.view(torch.uint8)
+ b_bytes = b_ref.view(torch.uint8)
+
+ nvfp4_gemv_module.nvfp4_gemv_batched_cuda(
+ a_bytes,
+ b_bytes,
+ sfa,
+ sfb,
+ c_ref,
+ M, K, L
+ )
+
+ return c_ref
scrolls · 457 diff lines total

Best evidence level for this revision: reported

JSON