Skip to content
KernelIndex
Search⌘K

submission 82459

mdouglas · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_cuda_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-82459?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
39.2µs
#219 of 678
2025-11-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:788607500ada4afc7425db45c71f6ad4b659e6af4d47a64ae083ed13a97e72ea
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[];

Kernel source

submission_cuda_v1.py591 lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import torch

# 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>
#include <cooperative_groups.h>

using namespace cooperative_groups;

// 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 multiple M rows, each warp handles one or more L batches.
// Inputs have native PyTorch strides from .permute() - K dimension has stride 1 (contiguous).
// Template parameters:
//   UseNestedLoops: true for large K (better coalescing), false for small K (better MLP)
//   LPerWarp: number of L batches each warp processes (2 for L=4,8)
template<bool UseNestedLoops, int LPerWarp>
__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;

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

    // Cooperatively load all L B vectors into shared memory
    // Strategy depends on K size:
    // - Large K (nested loops): exploit K-contiguity for coalesced global loads
    // - Small K (flat loops): better memory-level parallelism across L
    // B in memory: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half
    // We only need n=0 (the actual vector, rest is padding)
    if constexpr (UseNestedLoops) {
        // Large K: nested loops for coalesced loads
        for (int l = 0; l < L; l++) {
            for (int k = tid; k < K_half; k += blockDim.x) {
                sb[l * K_half + k] = b[k + l * N_padded * K_half];
            }
        }
    } else {
        // Small K: flat loops for better MLP across L
        for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
            int k = kl / L;
            int l = kl % L;
            sb[l * K_half + k] = b[k + l * N_padded * K_half];
        }
    }

    // Cooperatively load scale factors for B
    if constexpr (UseNestedLoops) {
        // Large K: nested loops
        for (int l = 0; l < L; l++) {
            for (int k = tid; k < K_div_16; k += blockDim.x) {
                ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
            }
        }
    } else {
        // Small K: flat loops
        for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
            int k = kl / L;
            int l = kl % L;
            ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
        }
    }

    __syncthreads();

    // Parallelize across L dimension with LPerWarp warps handling multiple L batches
    const int WARPS_PER_BLOCK = blockDim.x / 32;
    const int WARPS_PER_M_ROW = L / LPerWarp;  // Fewer warps when each handles multiple L
    const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;

    int m_base = blockIdx.x * M_ROWS_PER_BLOCK;

    // Which M row and which L batch group?
    int m_local = warp_id / WARPS_PER_M_ROW;
    int l_group = warp_id % WARPS_PER_M_ROW;
    int m = m_base + m_local;

    if (m >= M) return;

    // Each warp processes LPerWarp consecutive L batches
    int l_base = l_group * LPerWarp;

    // Process LPerWarp L batches per warp
    for (int l_offset = 0; l_offset < LPerWarp; l_offset++) {
        int l = l_base + l_offset;
        if (l >= L) break;

        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 2 scale blocks per iteration for better ILP (16 bytes with uint4)
        int num_scale_pairs = K_div_16 / 2;
        for (int scale_pair = lane; scale_pair < num_scale_pairs; scale_pair += 32) {
            int scale_block_0 = scale_pair * 2;

            // Vectorized loads and direct fp8x2 → half2 conversion
            __nv_fp8x2_e4m3 scale_a_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&sfa[sfa_base + scale_block_0]);
            __nv_fp8x2_e4m3 scale_b_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&ssfb_row[scale_block_0]);

            // Direct conversion to half2 and SIMD multiplication (compute both scales at once!)
            __half2 scales_a = static_cast<__half2>(scale_a_pair);
            __half2 scales_b = static_cast<__half2>(scale_b_pair);
            __half2 combined_scales = __hmul2(scales_a, scales_b);  // SIMD: both scales in one instruction

            // Broadcast each scale to half2 for use in compute loop
            __half2 scale2_0 = __half2half2(combined_scales.x);
            __half2 scale2_1 = __half2half2(combined_scales.y);

            // Load 16 bytes at once using uint4
            int k_byte_base = scale_block_0 * 8;

            // A matrix: streaming access, use .cg cache hint (L2 only)
            uint4 a_data;
            const uint32_t* a_ptr = reinterpret_cast<const uint32_t*>(&a[a_base + k_byte_base]);
            asm volatile("ld.global.cg.v4.u32 {%0,%1,%2,%3}, [%4];"
                : "=r"(a_data.x), "=r"(a_data.y), "=r"(a_data.z), "=r"(a_data.w)
                : "l"(a_ptr));

            // B from shared memory: contiguous access
            const uint4 b_data = *reinterpret_cast<const uint4*>(&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 first 8 bytes with scale_0 - use FMA for efficiency
            __half2 local_sum_0 = __float2half2_rn(0.0f);
            #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 product = __hmul2(a_vals, b_vals);
                local_sum_0 = __hfma2(product, scale2_0, local_sum_0);  // FMA: product * scale + sum
            }
            sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));

            // Process second 8 bytes with scale_1 - use FMA for efficiency
            __half2 local_sum_1 = __float2half2_rn(0.0f);
            #pragma unroll
            for (int i = 8; i < 16; 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 product = __hmul2(a_vals, b_vals);
                local_sum_1 = __hfma2(product, scale2_1, local_sum_1);  // FMA: product * scale + sum
            }
            sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
        }

        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
) {
    // 4 warps per M row - more iterations per thread for better ILP and latency hiding
    // 8 M rows per block to maximize work per block
    const int WARPS_PER_M_ROW = 4;
    const int M_ROWS_PER_BLOCK = 8;

    // Shared memory: B vector, scale factors B, and warp partial sums
    extern __shared__ uint8_t smem[];
    uint8_t* sb = smem;
    __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
    float* warp_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 with streaming cache hint
    // Use .cs (cache streaming, evict first) since B is loaded once per block
    int num_vec_loads = K_half / 16;
    for (int i = tid; i < num_vec_loads; i += blockDim.x) {
        uint4 data;
        const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(&b[i * 16]);
        asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
            : "=r"(data.x), "=r"(data.y), "=r"(data.z), "=r"(data.w)
            : "l"(b_ptr));
        reinterpret_cast<uint4*>(sb)[i] = data;
    }
    // Tail loop for B: handles remaining bytes when K_half not divisible by 16
    // K is always divisible by 64 per task spec, so K_half divisible by 32, usually by 16
    // This loop is unlikely to execute for valid inputs, commented out for clarity
    // int vec_bytes = num_vec_loads * 16;
    // for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
    //     unsigned int data;
    //     asm volatile("ld.global.cs.u8 %0, [%1];" : "=r"(data) : "l"(&b[i]));
    //     sb[i] = data;
    // }

    // Cooperatively load scale factors for B with streaming cache hint
    int num_vec_loads_sfb = K_div_16 / 16;
    for (int i = tid; i < num_vec_loads_sfb; i += blockDim.x) {
        uint4 data;
        const uint32_t* sfb_ptr = reinterpret_cast<const uint32_t*>(&sfb[i * 16]);
        asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
            : "=r"(data.x), "=r"(data.y), "=r"(data.z), "=r"(data.w)
            : "l"(sfb_ptr));
        reinterpret_cast<uint4*>(reinterpret_cast<uint8_t*>(ssfb))[i] = data;
    }
    // Tail loop for sfb: unlikely to execute for valid inputs
    // int vec_bytes_sfb = num_vec_loads_sfb * 16;
    // for (int i = tid + vec_bytes_sfb; i < K_div_16; i += blockDim.x) {
    //     unsigned int data;
    //     asm volatile("ld.global.cs.u8 %0, [%1];" : "=r"(data) : "l"(&sfb[i]));
    //     ssfb[i].__x = data;
    // }

    __syncthreads();

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

    // Which M row and K chunk does this warp handle?
    int m_local = warp_id / WARPS_PER_M_ROW;
    int m = m_base + m_local;

    if (m >= M) return;

    int warp_in_m_group = warp_id % WARPS_PER_M_ROW;

    // Process 2 scale blocks per iteration for better ILP
    int scale_pairs_per_warp = (K_div_16 / WARPS_PER_M_ROW) / 2;  // 1024 / 8 / 2 = 64 scale pairs per warp
    int scale_pair_start = warp_in_m_group * scale_pairs_per_warp;
    int scale_pair_end = scale_pair_start + scale_pairs_per_warp;

    float sum = 0.0f;

    // Loop over scale pairs - each iteration processes 16 bytes (2 scale blocks)
    // With 8 warps: 64 scale pairs / 32 threads = 2 iterations per thread
    for (int scale_pair = scale_pair_start + lane; scale_pair < scale_pair_end; scale_pair += 32) {
        int scale_block_0 = scale_pair * 2;

        // Load scale factors (sfa from global with constant cache, sfb from shared memory)
        __nv_fp8x2_e4m3 scale_a_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&sfa[m * K_div_16 + scale_block_0]);
        __nv_fp8x2_e4m3 scale_b_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&ssfb[scale_block_0]);

        // Direct conversion to half2 and SIMD multiplication (compute both scales at once!)
        __half2 scales_a = static_cast<__half2>(scale_a_pair);
        __half2 scales_b = static_cast<__half2>(scale_b_pair);
        __half2 combined_scales = __hmul2(scales_a, scales_b);  // SIMD: both scales in one instruction

        // Broadcast each scale to half2 for use in compute loop
        __half2 scale2_0 = __half2half2(combined_scales.x);
        __half2 scale2_1 = __half2half2(combined_scales.y);

        int k_byte_base = scale_block_0 * 8;

        // A matrix: streaming access, use .cg cache hint (L2 only)
        uint4 a_data;
        const uint32_t* a_ptr = reinterpret_cast<const uint32_t*>(&a[m * K_half + k_byte_base]);
        asm volatile("ld.global.cg.v4.u32 {%0,%1,%2,%3}, [%4];"
            : "=r"(a_data.x), "=r"(a_data.y), "=r"(a_data.z), "=r"(a_data.w)
            : "l"(a_ptr));

        // B from shared memory: normal load
        const uint4 b_data = *reinterpret_cast<const uint4*>(&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 first 8 bytes with scale_0 - use FMA for efficiency
        __half2 local_sum_0 = __float2half2_rn(0.0f);
        #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 product = __hmul2(a_vals, b_vals);
            local_sum_0 = __hfma2(product, scale2_0, local_sum_0);  // FMA: product * scale + sum
        }
        sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));

        // Process second 8 bytes with scale_1 - use FMA for efficiency
        __half2 local_sum_1 = __float2half2_rn(0.0f);
        #pragma unroll
        for (int i = 8; i < 16; 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 product = __hmul2(a_vals, b_vals);
            local_sum_1 = __hfma2(product, scale2_1, local_sum_1);  // FMA: product * scale + sum
        }
        sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
    }

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

    // Lane 0 of each warp writes its partial sum to shared memory
    if (lane == 0) {
        warp_sums[warp_id] = sum;
    }

    __syncthreads();

    // Final reduction: each of the first M_ROWS_PER_BLOCK threads reduces one M row
    if (tid < M_ROWS_PER_BLOCK) {
        int m_write = m_base + tid;
        if (m_write < M) {
            // Reduce WARPS_PER_M_ROW partial sums for this M row
            float final_sum = 0.0f;
            int warp_start = tid * WARPS_PER_M_ROW;
            #pragma unroll
            for (int w = 0; w < WARPS_PER_M_ROW; w++) {
                final_sum += warp_sums[warp_start + w];
            }
            c[m_write] = __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
) {
    // 4 warps per M row, 8 M rows per block
    const int WARPS_PER_M_ROW = 4;
    const int M_ROWS_PER_BLOCK = 8;
    const int WARPS_PER_BLOCK = WARPS_PER_M_ROW * M_ROWS_PER_BLOCK;  // 32
    const int threads = WARPS_PER_BLOCK * 32;  // 1024 threads
    const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

    // Shared memory: B vector + sfb + warp 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
) {
    const int WARPS_PER_BLOCK = 32;
    const int threads = WARPS_PER_BLOCK * 32;  // 1024 threads

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

    const int K_half = K / 2;
    const int K_div_16 = K / 16;
    const int smem_size = L * K_half + L * K_div_16;

    // Choose loading strategy based on K size:
    // Small K (<=2048): use flat loops for better memory-level parallelism
    // Large K (>2048): use nested loops for better coalescing
    const bool use_nested_loops = (K > 2048);

    // Dispatch based on L to optimize LPerWarp for ILP and shared memory amortization
    if (L == 8) {
        // L=8: 2 L batches per warp for better ILP and shared memory amortization
        // 8 M rows per block, 512 blocks total
        const int L_PER_WARP = 2;
        const int WARPS_PER_M_ROW = L / L_PER_WARP;  // 4
        const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;  // 8
        const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

        if (use_nested_loops) {
            nvfp4_gemv_batched_kernel<true, 2><<<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);
        } else {
            nvfp4_gemv_batched_kernel<false, 2><<<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);
        }
    } else if (L == 4) {
        // L=4: 2 L batches per warp, 16 M rows per block
        const int L_PER_WARP = 2;
        const int WARPS_PER_M_ROW = L / L_PER_WARP;  // 2
        const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;  // 16
        const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

        if (use_nested_loops) {
            nvfp4_gemv_batched_kernel<true, 2><<<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);
        } else {
            nvfp4_gemv_batched_kernel<false, 2><<<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);
        }
    } else {
        // Default: 1 L per warp
        const int WARPS_PER_M_ROW = L;
        const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;
        const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;

        if (use_nested_loops) {
            nvfp4_gemv_batched_kernel<true, 1><<<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);
        } else {
            nvfp4_gemv_batched_kernel<false, 1><<<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());
}

// Dispatch function - handles L=1 and batched L cases
void nvfp4_gemv_dispatch_cuda(
    torch::Tensor a_ref,      // [M, K//2, L] FP4
    torch::Tensor b_ref,      // [128, K//2, L] FP4
    torch::Tensor sfa,        // [M, K//16, L] FP8
    torch::Tensor sfb,        // [128, K//16, L] FP8
    torch::Tensor c_ref       // [M, 1, L] FP16
) {
    int M = a_ref.size(0);
    int K_half = a_ref.size(1);
    int L = a_ref.size(2);
    int K = K_half * 2;

    if (L == 1) {
        // L=1: Extract L=0 slice (all operations are views, no copying)
        auto a_slice = a_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});
        auto b_slice = b_ref.index({0, torch::indexing::Slice(), 0});
        auto sfa_slice = sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});
        auto sfb_slice = sfb.index({0, torch::indexing::Slice(), 0});

        auto a_bytes = a_slice.view(torch::kUInt8).contiguous();
        auto b_bytes = b_slice.view(torch::kUInt8).contiguous();

        nvfp4_gemv_cuda(
            a_bytes,
            b_bytes,
            sfa_slice.contiguous(),
            sfb_slice.contiguous(),
            c_ref,
            M, K
        );
    } else {
        // Other L: use batched kernel
        auto a_bytes = a_ref.view(torch::kUInt8);
        auto b_bytes = b_ref.view(torch::kUInt8);

        nvfp4_gemv_batched_cuda(
            a_bytes, b_bytes, sfa, sfb, c_ref,
            M, K, L
        );
    }
}
"""

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

# Compile the CUDA extension inline
nvfp4_gemv_module = load_inline(
    name='nvfp4_gemv',
    cpp_sources=[cpp_source],
    cuda_sources=[cuda_source],
    functions=['nvfp4_gemv_dispatch_cuda'],
    verbose=True,
    extra_cflags=[
        '-O3',
        '-std=c++17',
        '-march=native',
    ],
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '--extra-device-vectorization',
        '-arch=sm_100a',
        '-std=c++17',
        '-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:
    """
    Custom CUDA implementation of NVFP4 block-scaled GEMV.
    Dispatch logic now in C++ to minimize Python overhead.
    """
    a_ref, b_ref, sfa, sfb, _, _, c_ref = data

    # Single C++ call - dispatch logic handled in C++
    nvfp4_gemv_module.nvfp4_gemv_dispatch_cuda(
        a_ref,
        b_ref,
        sfa,
        sfb,
        c_ref
    )

    return c_ref
scrolls · 591 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 81959.

⋯ 11 unchanged lines
#include <cuda_fp8.h>
#include <ATen/cuda/Exceptions.h>
#include <ATen/cuda/CUDAContext.h>
+ #include <cooperative_groups.h>
+ using namespace cooperative_groups;
+
// 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
⋯ 9 unchanged lines
// Batched kernel: processes all L batches in one launch for better efficiency.
// Each block handles multiple M rows, each warp handles one or more L batches.
- // Inputs have native PyTorch strides from .permute() - K dimension has stride 1.
- // Template parameter UseNestedLoops: true for large K (avoid div/mod), false for small K (less overhead)
- // Template parameter LPerWarp: number of L batches each warp processes (1 for L=8, 2 for L=4)
+ // Inputs have native PyTorch strides from .permute() - K dimension has stride 1 (contiguous).
+ // Template parameters:
+ // UseNestedLoops: true for large K (better coalescing), false for small K (better MLP)
+ // LPerWarp: number of L batches each warp processes (2 for L=4,8)
template<bool UseNestedLoops, int LPerWarp>
__global__ void nvfp4_gemv_batched_kernel(
const uint8_t* __restrict__ a, // [M, K//2, L] with strides (K_half, 1, M*K_half)
⋯ 20 unchanged lines
const int N_padded = 128; // B is padded to 128 rows for torch._scaled_mm
// Cooperatively load all L B vectors into shared memory
- // Template specialization: compile-time branch selection based on K size
- // B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half
+ // Strategy depends on K size:
+ // - Large K (nested loops): exploit K-contiguity for coalesced global loads
+ // - Small K (flat loops): better memory-level parallelism across L
+ // B in memory: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half
// We only need n=0 (the actual vector, rest is padding)
if constexpr (UseNestedLoops) {
- // Large K: nested loops to avoid expensive div/mod
+ // Large K: nested loops for coalesced loads
for (int l = 0; l < L; l++) {
for (int k = tid; k < K_half; k += blockDim.x) {
sb[l * K_half + k] = b[k + l * N_padded * K_half];
}
}
} else {
- // Small K: flat loop with div/mod has less overhead
+ // Small K: flat loops for better MLP across L
for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
⋯ 10 unchanged lines
}
}
} else {
- // Small K: flat loop
+ // Small K: flat loops
for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
⋯ 37 unchanged lines
int num_scale_pairs = K_div_16 / 2;
for (int scale_pair = lane; scale_pair < num_scale_pairs; scale_pair += 32) {
int scale_block_0 = scale_pair * 2;
- int scale_block_1 = scale_block_0 + 1;
- // Load both scale factors
- half scale_a0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block_0].__x, __NV_E4M3).x);
- half scale_b0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block_0].__x, __NV_E4M3).x);
- half combined_scale0 = scale_a0 * scale_b0;
- __half2 scale2_0 = __half2half2(combined_scale0);
+ // Vectorized loads and direct fp8x2 → half2 conversion
+ __nv_fp8x2_e4m3 scale_a_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&sfa[sfa_base + scale_block_0]);
+ __nv_fp8x2_e4m3 scale_b_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&ssfb_row[scale_block_0]);
- half scale_a1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block_1].__x, __NV_E4M3).x);
- half scale_b1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block_1].__x, __NV_E4M3).x);
- half combined_scale1 = scale_a1 * scale_b1;
- __half2 scale2_1 = __half2half2(combined_scale1);
+ // Direct conversion to half2 and SIMD multiplication (compute both scales at once!)
+ __half2 scales_a = static_cast<__half2>(scale_a_pair);
+ __half2 scales_b = static_cast<__half2>(scale_b_pair);
+ __half2 combined_scales = __hmul2(scales_a, scales_b); // SIMD: both scales in one instruction
+ // Broadcast each scale to half2 for use in compute loop
+ __half2 scale2_0 = __half2half2(combined_scales.x);
+ __half2 scale2_1 = __half2half2(combined_scales.y);
+
// Load 16 bytes at once using uint4
int k_byte_base = scale_block_0 * 8;
⋯ 4 unchanged lines
: "=r"(a_data.x), "=r"(a_data.y), "=r"(a_data.z), "=r"(a_data.w)
: "l"(a_ptr));
- // B from shared memory: normal load
+ // B from shared memory: contiguous access
const uint4 b_data = *reinterpret_cast<const uint4*>(&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 first 8 bytes with scale_0
+ // Process first 8 bytes with scale_0 - use FMA for efficiency
__half2 local_sum_0 = __float2half2_rn(0.0f);
#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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_0);
- local_sum_0 = __hadd2(local_sum_0, scaled);
+ __half2 product = __hmul2(a_vals, b_vals);
+ local_sum_0 = __hfma2(product, scale2_0, local_sum_0); // FMA: product * scale + sum
}
sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));
- // Process second 8 bytes with scale_1
+ // Process second 8 bytes with scale_1 - use FMA for efficiency
__half2 local_sum_1 = __float2half2_rn(0.0f);
#pragma unroll
for (int i = 8; i < 16; 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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_1);
- local_sum_1 = __hadd2(local_sum_1, scaled);
+ __half2 product = __hmul2(a_vals, b_vals);
+ local_sum_1 = __hfma2(product, scale2_1, local_sum_1); // FMA: product * scale + sum
}
sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
}
⋯ 26 unchanged lines
extern __shared__ uint8_t smem[];
uint8_t* sb = smem;
__nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
- float* warp_sums = reinterpret_cast<float*>(ssfb + K/16); // 32 warp partial sums
+ float* warp_sums = reinterpret_cast<float*>(ssfb + K/16);
int tid = threadIdx.x;
int warp_id = tid / 32;
⋯ 2 unchanged lines
int K_half = K / 2;
int K_div_16 = K / 16;
- // Cooperatively load B vector into shared memory (vectorized)
+ // Cooperatively load B vector into shared memory with streaming cache hint
+ // Use .cs (cache streaming, evict first) since B is loaded once per block
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];
+ uint4 data;
+ const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(&b[i * 16]);
+ asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
+ : "=r"(data.x), "=r"(data.y), "=r"(data.z), "=r"(data.w)
+ : "l"(b_ptr));
+ reinterpret_cast<uint4*>(sb)[i] = data;
}
- int vec_bytes = num_vec_loads * 16;
- for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
- sb[i] = b[i];
- }
+ // Tail loop for B: handles remaining bytes when K_half not divisible by 16
+ // K is always divisible by 64 per task spec, so K_half divisible by 32, usually by 16
+ // This loop is unlikely to execute for valid inputs, commented out for clarity
+ // int vec_bytes = num_vec_loads * 16;
+ // for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
+ // unsigned int data;
+ // asm volatile("ld.global.cs.u8 %0, [%1];" : "=r"(data) : "l"(&b[i]));
+ // sb[i] = data;
+ // }
- // Cooperatively load scale factors for B
- for (int i = tid; i < K_div_16; i += blockDim.x) {
- ssfb[i] = sfb[i];
+ // Cooperatively load scale factors for B with streaming cache hint
+ int num_vec_loads_sfb = K_div_16 / 16;
+ for (int i = tid; i < num_vec_loads_sfb; i += blockDim.x) {
+ uint4 data;
+ const uint32_t* sfb_ptr = reinterpret_cast<const uint32_t*>(&sfb[i * 16]);
+ asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
+ : "=r"(data.x), "=r"(data.y), "=r"(data.z), "=r"(data.w)
+ : "l"(sfb_ptr));
+ reinterpret_cast<uint4*>(reinterpret_cast<uint8_t*>(ssfb))[i] = data;
}
+ // Tail loop for sfb: unlikely to execute for valid inputs
+ // int vec_bytes_sfb = num_vec_loads_sfb * 16;
+ // for (int i = tid + vec_bytes_sfb; i < K_div_16; i += blockDim.x) {
+ // unsigned int data;
+ // asm volatile("ld.global.cs.u8 %0, [%1];" : "=r"(data) : "l"(&sfb[i]));
+ // ssfb[i].__x = data;
+ // }
__syncthreads();
- // Each block processes 2 M rows
+ // Each block processes 8 M rows
int m_base = blockIdx.x * M_ROWS_PER_BLOCK;
// Which M row and K chunk does this warp handle?
⋯ 15 unchanged lines
// With 8 warps: 64 scale pairs / 32 threads = 2 iterations per thread
for (int scale_pair = scale_pair_start + lane; scale_pair < scale_pair_end; scale_pair += 32) {
int scale_block_0 = scale_pair * 2;
- int scale_block_1 = scale_block_0 + 1;
- // Load both scale factors
- half scale_a0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block_0].__x, __NV_E4M3).x);
- half scale_b0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block_0].__x, __NV_E4M3).x);
- half combined_scale0 = scale_a0 * scale_b0;
- __half2 scale2_0 = __half2half2(combined_scale0);
+ // Load scale factors (sfa from global with constant cache, sfb from shared memory)
+ __nv_fp8x2_e4m3 scale_a_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&sfa[m * K_div_16 + scale_block_0]);
+ __nv_fp8x2_e4m3 scale_b_pair = *reinterpret_cast<const __nv_fp8x2_e4m3*>(&ssfb[scale_block_0]);
- half scale_a1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block_1].__x, __NV_E4M3).x);
- half scale_b1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block_1].__x, __NV_E4M3).x);
- half combined_scale1 = scale_a1 * scale_b1;
- __half2 scale2_1 = __half2half2(combined_scale1);
+ // Direct conversion to half2 and SIMD multiplication (compute both scales at once!)
+ __half2 scales_a = static_cast<__half2>(scale_a_pair);
+ __half2 scales_b = static_cast<__half2>(scale_b_pair);
+ __half2 combined_scales = __hmul2(scales_a, scales_b); // SIMD: both scales in one instruction
+ // Broadcast each scale to half2 for use in compute loop
+ __half2 scale2_0 = __half2half2(combined_scales.x);
+ __half2 scale2_1 = __half2half2(combined_scales.y);
+
int k_byte_base = scale_block_0 * 8;
// A matrix: streaming access, use .cg cache hint (L2 only)
⋯ 9 unchanged lines
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 first 8 bytes with scale_0 - accumulate in half2 first
+ // Process first 8 bytes with scale_0 - use FMA for efficiency
__half2 local_sum_0 = __float2half2_rn(0.0f);
#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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_0);
- local_sum_0 = __hadd2(local_sum_0, scaled);
+ __half2 product = __hmul2(a_vals, b_vals);
+ local_sum_0 = __hfma2(product, scale2_0, local_sum_0); // FMA: product * scale + sum
}
sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));
- // Process second 8 bytes with scale_1 - accumulate in half2 first
+ // Process second 8 bytes with scale_1 - use FMA for efficiency
__half2 local_sum_1 = __float2half2_rn(0.0f);
#pragma unroll
for (int i = 8; i < 16; 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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_1);
- local_sum_1 = __hadd2(local_sum_1, scaled);
+ __half2 product = __hmul2(a_vals, b_vals);
+ local_sum_1 = __hfma2(product, scale2_1, local_sum_1); // FMA: product * scale + sum
}
sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
}
⋯ 79 unchanged lines
const int K_div_16 = K / 16;
const int smem_size = L * K_half + L * K_div_16;
- // For L=4 with small K: each warp handles 2 L batches for better ILP
- // For L=8: each warp handles 1 L batch
- if (L == 4 && K < 4096) {
- // L=4, small K: 2 L batches per warp, 16 M rows per block
+ // Choose loading strategy based on K size:
+ // Small K (<=2048): use flat loops for better memory-level parallelism
+ // Large K (>2048): use nested loops for better coalescing
+ const bool use_nested_loops = (K > 2048);
+
+ // Dispatch based on L to optimize LPerWarp for ILP and shared memory amortization
+ if (L == 8) {
+ // L=8: 2 L batches per warp for better ILP and shared memory amortization
+ // 8 M rows per block, 512 blocks total
const int L_PER_WARP = 2;
+ const int WARPS_PER_M_ROW = L / L_PER_WARP; // 4
+ const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // 8
+ const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
+
+ if (use_nested_loops) {
+ nvfp4_gemv_batched_kernel<true, 2><<<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);
+ } else {
+ nvfp4_gemv_batched_kernel<false, 2><<<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);
+ }
+ } else if (L == 4) {
+ // L=4: 2 L batches per warp, 16 M rows per block
+ const int L_PER_WARP = 2;
const int WARPS_PER_M_ROW = L / L_PER_WARP; // 2
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // 16
const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
- nvfp4_gemv_batched_kernel<false, 2><<<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
- );
- } else if (K >= 4096) {
- // Large K: use nested loops, 1 L per warp
- const int WARPS_PER_M_ROW = L;
- const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;
- const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
-
- nvfp4_gemv_batched_kernel<true, 1><<<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
- );
+ if (use_nested_loops) {
+ nvfp4_gemv_batched_kernel<true, 2><<<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);
+ } else {
+ nvfp4_gemv_batched_kernel<false, 2><<<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);
+ }
} else {
- // Default: small K, 1 L per warp
+ // Default: 1 L per warp
const int WARPS_PER_M_ROW = L;
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;
const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
- nvfp4_gemv_batched_kernel<false, 1><<<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
- );
+ if (use_nested_loops) {
+ nvfp4_gemv_batched_kernel<true, 1><<<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);
+ } else {
+ nvfp4_gemv_batched_kernel<false, 1><<<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());
}
- // Dispatch function - handles L=1 extraction and L=8 splitting in C++
+ // Dispatch function - handles L=1 and batched L cases
void nvfp4_gemv_dispatch_cuda(
torch::Tensor a_ref, // [M, K//2, L] FP4
torch::Tensor b_ref, // [128, K//2, L] FP4
⋯ 66 unchanged lines
'-arch=sm_100a',
'-std=c++17',
'-Xptxas', '-v',
+ '-lineinfo',
'-U__CUDA_NO_HALF_OPERATORS__', # Enable half operators
'-U__CUDA_NO_HALF_CONVERSIONS__', # Enable half conversions
],
⋯ 15 unchanged lines
c_ref
)
- return c_ref
+ return c_ref
No newline at end of file
scrolls · 391 diff lines total

Best evidence level for this revision: reported

JSON