Skip to content
KernelIndex
Search⌘K

submission 105811

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:26a50c5a6b19002a1c2f92118744c4dee7914464e8e189a516a815117677860c
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

async-copyasm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
fp41. Native FP4/FP8 conversion functions
fp8__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
shared-memory__shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
vector-width = uint4uint4 a_vec = *reinterpret_cast<const uint4*>(A + row * k_packed + k_global);

Kernel source

submission.py413 lines
"""
Enhanced GEMV based on our working baseline with selective improvements:
1. Native FP4/FP8 conversion functions 
2. Double buffering for B and SFB with async copy
3. Keep our thread block structure (32x32 = 1024 threads)
4. Keep our per-thread work (16 bytes per iteration)
5. SM100a architecture flag for Blackwell
"""

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cstdint>

#define THREADS_PER_M 32
#define THREADS_PER_K 32
#define BLOCK_SIZE (THREADS_PER_M * THREADS_PER_K)
#define SF_VEC_SIZE 16
#define K_TILE_SIZE 512  // 32 threads × 16 bytes
#define SF_TILE_SIZE 64  // K_TILE_SIZE / 8

// Native FP4 E2M1 to half2 conversion (from reference)
__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
    __half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
        static_cast<__nv_fp4x2_storage_t>(byte),
        __NV_E2M1
    );
    return *reinterpret_cast<__half2*>(&raw);
}

// Native FP8 E4M3 to float conversion (from reference)
__device__ __forceinline__ float decode_fp8(uint8_t byte) {
    __nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
    __half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
    return __half2float(__ushort_as_half(raw.x));
}

// Async copy macros
#define ASYNC_COPY_16(dst, src) \
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))

#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
#define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")

// Kernel for l == 1 with double-buffered B
__global__ __launch_bounds__(BLOCK_SIZE, 2)
void nvfp4_gemv_kernel_single(
    half* __restrict__ C,
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    int m, int k_packed, int k_sf, int n_pad
) {
    // Double buffer for B (shared across all rows)
    __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
    __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
    
    const int tidx = threadIdx.x;
    const int tidy = threadIdx.y;
    const int tid = tidy * THREADS_PER_K + tidx;
    const int row = blockIdx.x * THREADS_PER_M + tidy;
    const bool valid = (row < m);
    
    const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
    
    float sum = 0.0f;
    
    // Load first tile
    if (num_tiles > 0) {
        if (tid < 32) {
            int k_off = tid * 16;
            if (k_off < k_packed) {
                uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
                ASYNC_COPY_16(dst, B + k_off);
            }
        }
        if (tid < 4) {
            int sf_off = tid * 16;
            if (sf_off < k_sf) {
                uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
                ASYNC_COPY_16(dst, SFB + sf_off);
            }
        }
        ASYNC_COMMIT();
    }
    
    for (int tile = 0; tile < num_tiles; tile++) {
        int curr = tile & 1;
        int next = (tile + 1) & 1;
        int next_tile = tile + 1;
        
        // Prefetch next tile
        if (next_tile < num_tiles) {
            int k_base_next = next_tile * K_TILE_SIZE;
            if (tid < 32) {
                int k_off = k_base_next + tid * 16;
                if (k_off < k_packed) {
                    uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
                    ASYNC_COPY_16(dst, B + k_off);
                }
            }
            if (tid < 4) {
                int sf_off = k_base_next / 8 + tid * 16;
                if (sf_off < k_sf) {
                    uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
                    ASYNC_COPY_16(dst, SFB + sf_off);
                }
            }
            ASYNC_COMMIT();
        }
        
        // Wait for current tile
        ASYNC_WAIT_ALL();
        __syncthreads();
        
        if (valid) {
            int k_base = tile * K_TILE_SIZE;
            int k_local = tidx * 16;
            int k_global = k_base + k_local;
            
            if (k_global + 16 <= k_packed) {
                // Load A directly (per-row, can't share)
                uint4 a_vec = *reinterpret_cast<const uint4*>(A + row * k_packed + k_global);
                // Load B from shared memory
                uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
                
                const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
                const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
                
                int sf_local = k_local / 8;
                int sf_global = k_global / 8;
                
                float sfa0 = decode_fp8(SFA[row * k_sf + sf_global]);
                float sfa1 = decode_fp8(SFA[row * k_sf + sf_global + 1]);
                float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
                float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
                float scale0 = sfa0 * sfb0;
                float scale1 = sfa1 * sfb1;
                
                __half2 acc0 = __float2half2_rn(0.0f);
                __half2 acc1 = __float2half2_rn(0.0f);
                
                #pragma unroll
                for (int i = 0; i < 8; i++) {
                    acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
                }
                sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
                
                #pragma unroll
                for (int i = 8; i < 16; i++) {
                    acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
                }
                sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
            }
        }
        
        __syncthreads();
    }
    
    if (!valid) return;
    
    // Warp shuffle reduction
    #pragma unroll
    for (int offset = 16; offset > 0; offset /= 2) {
        sum += __shfl_down_sync(0xffffffff, sum, offset);
    }
    
    if (tidx == 0) {
        C[row] = __float2half(sum);
    }
}

// Kernel for l > 1 with double-buffered B
__global__ __launch_bounds__(BLOCK_SIZE, 2)
void nvfp4_gemv_kernel_batched(
    half* __restrict__ C,
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    int m, int k_packed, int k_sf, int l, int n_pad
) {
    __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
    __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
    
    const int tidx = threadIdx.x;
    const int tidy = threadIdx.y;
    const int tid = tidy * THREADS_PER_K + tidx;
    const int row = blockIdx.x * THREADS_PER_M + tidy;
    const int batch = blockIdx.z;
    const bool valid = (row < m);
    
    // Batch offsets
    const size_t a_batch = batch * (size_t)(m * k_packed);
    const size_t b_batch = batch * (size_t)(n_pad * k_packed);
    const size_t sfa_batch = batch * (size_t)(m * k_sf);
    const size_t sfb_batch = batch * (size_t)(n_pad * k_sf);
    
    const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
    
    float sum = 0.0f;
    
    // Load first tile
    if (num_tiles > 0) {
        if (tid < 32) {
            int k_off = tid * 16;
            if (k_off < k_packed) {
                uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
                ASYNC_COPY_16(dst, B + b_batch + k_off);
            }
        }
        if (tid < 4) {
            int sf_off = tid * 16;
            if (sf_off < k_sf) {
                uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
                ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
            }
        }
        ASYNC_COMMIT();
    }
    
    for (int tile = 0; tile < num_tiles; tile++) {
        int curr = tile & 1;
        int next = (tile + 1) & 1;
        int next_tile = tile + 1;
        
        // Prefetch next tile
        if (next_tile < num_tiles) {
            int k_base_next = next_tile * K_TILE_SIZE;
            if (tid < 32) {
                int k_off = k_base_next + tid * 16;
                if (k_off < k_packed) {
                    uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
                    ASYNC_COPY_16(dst, B + b_batch + k_off);
                }
            }
            if (tid < 4) {
                int sf_off = k_base_next / 8 + tid * 16;
                if (sf_off < k_sf) {
                    uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
                    ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
                }
            }
            ASYNC_COMMIT();
        }
        
        ASYNC_WAIT_ALL();
        __syncthreads();
        
        if (valid) {
            int k_base = tile * K_TILE_SIZE;
            int k_local = tidx * 16;
            int k_global = k_base + k_local;
            
            if (k_global + 16 <= k_packed) {
                uint4 a_vec = *reinterpret_cast<const uint4*>(A + a_batch + row * k_packed + k_global);
                uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
                
                const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
                const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
                
                int sf_local = k_local / 8;
                int sf_global = k_global / 8;
                
                float sfa0 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global]);
                float sfa1 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global + 1]);
                float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
                float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
                float scale0 = sfa0 * sfb0;
                float scale1 = sfa1 * sfb1;
                
                __half2 acc0 = __float2half2_rn(0.0f);
                __half2 acc1 = __float2half2_rn(0.0f);
                
                #pragma unroll
                for (int i = 0; i < 8; i++) {
                    acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
                }
                sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
                
                #pragma unroll
                for (int i = 8; i < 16; i++) {
                    acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
                }
                sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
            }
        }
        
        __syncthreads();
    }
    
    if (!valid) return;
    
    #pragma unroll
    for (int offset = 16; offset > 0; offset /= 2) {
        sum += __shfl_down_sync(0xffffffff, sum, offset);
    }
    
    if (tidx == 0) {
        C[batch * m + row] = __float2half(sum);
    }
}

void run_nvfp4_gemv(
    torch::Tensor C,
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    int m, int k, int l, int n_pad
) {
    int k_packed = k / 2;
    int k_sf = k / SF_VEC_SIZE;
    
    half* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
    const uint8_t* a_ptr = A.data_ptr<uint8_t>();
    const uint8_t* b_ptr = B.data_ptr<uint8_t>();
    const uint8_t* sfa_ptr = SFA.data_ptr<uint8_t>();
    const uint8_t* sfb_ptr = SFB.data_ptr<uint8_t>();
    
    int blocks_m = (m + THREADS_PER_M - 1) / THREADS_PER_M;
    dim3 block(THREADS_PER_K, THREADS_PER_M, 1);

    if (l == 1) {
        dim3 grid(blocks_m, 1, 1);
        nvfp4_gemv_kernel_single<<<grid, block>>>(
            c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
            m, k_packed, k_sf, n_pad
        );
    } else {
        dim3 grid(blocks_m, 1, l);
        nvfp4_gemv_kernel_batched<<<grid, block>>>(
            c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
            m, k_packed, k_sf, l, n_pad
        );
    }
}

#undef ASYNC_COPY_16
#undef ASYNC_COMMIT
#undef ASYNC_WAIT_ALL
'''

CPP_SRC = r'''
#include <torch/extension.h>

void run_nvfp4_gemv(
    torch::Tensor C,
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    int m, int k, int l, int n_pad
);
'''

_cuda_module = None

def get_cuda_module():
    global _cuda_module
    if _cuda_module is None:
        _cuda_module = load_inline(
            name='nvfp4_gemv_enhanced_v1',
            cpp_sources=CPP_SRC,
            cuda_sources=CUDA_SRC,
            functions=['run_nvfp4_gemv'],
            extra_cuda_cflags=[
                '-O3', 
                '--use_fast_math',
                '-std=c++17',
                '-gencode=arch=compute_100a,code=sm_100a'
            ],
            verbose=False,
        )
    return _cuda_module


def custom_kernel(data: input_t) -> output_t:
    a, b, sfa_ref, sfb_ref, _, _, c = data
    
    module = get_cuda_module()
    
    m, k_packed, l = a.shape
    k = k_packed * 2
    n_pad = b.shape[0]
    
    if not sfa_ref.is_cuda:
        sfa_ref = sfa_ref.to(a.device)
    if not sfb_ref.is_cuda:
        sfb_ref = sfb_ref.to(a.device)
    
    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)
    sfa_uint8 = sfa_ref.view(torch.uint8)
    sfb_uint8 = sfb_ref.view(torch.uint8)
    
    c_out = c.squeeze(1)
    
    module.run_nvfp4_gemv(
        c_out, a_uint8, b_uint8, sfa_uint8, sfb_uint8,
        m, k, l, n_pad
    )
    
    return c
scrolls · 413 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 103911.

"""
- CuTe DSL implementation of NVFP4 block-scaled GEMV.
-
- This is a simplified version that follows the same pattern as submission_cute.py
- but with cleaner structure. The kernel processes all batches in a single launch.
+ Enhanced GEMV based on our working baseline with selective improvements:
+ 1. Native FP4/FP8 conversion functions
+ 2. Double buffering for B and SFB with async copy
+ 3. Keep our thread block structure (32x32 = 1024 threads)
+ 4. Keep our per-thread work (16 bytes per iteration)
+ 5. SM100a architecture flag for Blackwell
"""
import torch
+ from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
- import cutlass
- import cutlass.cute as cute
- from cutlass.cute.runtime import make_ptr
- import cutlass.utils.blockscaled_layout as blockscaled_utils
+ CUDA_SRC = r'''
+ #include <torch/extension.h>
+ #include <cuda_runtime.h>
+ #include <cuda_fp16.h>
+ #include <cuda_fp4.h>
+ #include <cuda_fp8.h>
+ #include <cstdint>
- from cutlass import Float32
- from cutlass.cutlass_dsl import T, dsl_user_op
- from cutlass._mlir.dialects import nvvm
+ #define THREADS_PER_M 32
+ #define THREADS_PER_K 32
+ #define BLOCK_SIZE (THREADS_PER_M * THREADS_PER_K)
+ #define SF_VEC_SIZE 16
+ #define K_TILE_SIZE 512 // 32 threads × 16 bytes
+ #define SF_TILE_SIZE 64 // K_TILE_SIZE / 8
- # Kernel configuration parameters
- ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and B
- sf_dtype = cutlass.Float8E4M3FN # FP8 data type for scale factors
- c_dtype = cutlass.Float16 # FP16 output type
- accum_dtype = cutlass.Float32
- sf_vec_size = 16 # Scale factor block size (16 elements share one scale)
+ // Native FP4 E2M1 to half2 conversion (from reference)
+ __device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
+ __half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
+ static_cast<__nv_fp4x2_storage_t>(byte),
+ __NV_E2M1
+ );
+ return *reinterpret_cast<__half2*>(&raw);
+ }
- # Thread block configuration
- threads_per_m = 32
- threads_per_k = 4
- blk_k = 256 # K tile size
+ // Native FP8 E4M3 to float conversion (from reference)
+ __device__ __forceinline__ float decode_fp8(uint8_t byte) {
+ __nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
+ __half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
+ return __half2float(__ushort_as_half(raw.x));
+ }
- # Tile sizes for the mainloop
- mma_tiler_mnk = (threads_per_m, 1, blk_k)
+ // Async copy macros
+ #define ASYNC_COPY_16(dst, src) \
+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
+ #define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
+ #define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")
- def ceil_div(a, b):
- return (a + b - 1) // b
-
-
- @dsl_user_op
- def atomic_add_fp32(a: float | Float32, gmem_ptr: cute.Pointer, *, loc=None, ip=None) -> None:
- nvvm.atomicrmw(
- res=T.f32(), op=nvvm.AtomicOpKind.FADD, ptr=gmem_ptr.llvm_ptr, a=Float32(a).ir_value()
- )
-
-
- @dsl_user_op
- def elem_pointer(x: cute.Tensor, coord: cute.Coord, *, loc=None, ip=None) -> cute.Pointer:
- return x.iterator + cute.crd2idx(coord, x.layout, loc=loc, ip=ip)
-
-
- @cute.jit
- def scalar_to_ssa(a: cute.Numeric, dtype) -> cute.TensorSSA:
- """Convert a scalar to a cute TensorSSA of shape (1,) and given dtype."""
- vec = cute.make_fragment(1, dtype)
- vec[0] = a
- return vec.load()
-
-
- @cute.kernel
- def gemv_kernel(
- mA_mkl: cute.Tensor,
- mB_nkl: cute.Tensor,
- mSFA_mkl: cute.Tensor,
- mSFB_nkl: cute.Tensor,
- mC_mnl: cute.Tensor,
- ):
- """
- Block-scaled GEMV kernel.
+ // Kernel for l == 1 with double-buffered B
+ __global__ __launch_bounds__(BLOCK_SIZE, 2)
+ void nvfp4_gemv_kernel_single(
+ half* __restrict__ C,
+ const uint8_t* __restrict__ A,
+ const uint8_t* __restrict__ B,
+ const uint8_t* __restrict__ SFA,
+ const uint8_t* __restrict__ SFB,
+ int m, int k_packed, int k_sf, int n_pad
+ ) {
+ // Double buffer for B (shared across all rows)
+ __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
+ __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
- Computes: C[m, 1, l] = sum_k(A[m, k, l] * SFA[m, k, l] * B[n, k, l] * SFB[n, k, l])
+ const int tidx = threadIdx.x;
+ const int tidy = threadIdx.y;
+ const int tid = tidy * THREADS_PER_K + tidx;
+ const int row = blockIdx.x * THREADS_PER_M + tidy;
+ const bool valid = (row < m);
- Grid: (ceil(m/threads_per_m), 1, l)
- Block: (threads_per_m, threads_per_k, 1)
- """
- bidx, bidy, bidz = cute.arch.block_idx()
- tidx, tidy, _ = cute.arch.thread_idx()
+ const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
- # Extract tiles for A and its scale factors
- gA_mkl = cute.local_tile(
- mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
- gSFA_mkl = cute.local_tile(
- mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
+ float sum = 0.0f;
- # Extract tiles for B and its scale factors
- gB_nkl = cute.local_tile(
- mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
- gSFB_nkl = cute.local_tile(
- mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
+ // Load first tile
+ if (num_tiles > 0) {
+ if (tid < 32) {
+ int k_off = tid * 16;
+ if (k_off < k_packed) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
+ ASYNC_COPY_16(dst, B + k_off);
+ }
+ }
+ if (tid < 4) {
+ int sf_off = tid * 16;
+ if (sf_off < k_sf) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
+ ASYNC_COPY_16(dst, SFB + sf_off);
+ }
+ }
+ ASYNC_COMMIT();
+ }
- # Extract tiles for output C
- gC_mnl = cute.local_tile(
- mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
- )
-
- # Select output element for this thread
- tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
- tCgC = cute.make_tensor(tCgC.iterator, 1)
-
- # Initialize accumulator in FP32
- res = cute.zeros_like(tCgC, accum_dtype)
-
- # Shared memory for reduction across K dimension
- allocator = cutlass.utils.SmemAllocator()
- smem_layout = cute.make_layout(threads_per_m)
- shared_res = allocator.allocate_tensor(
- element_type=cutlass.Float32, layout=smem_layout
- )
-
- # Initialize shared memory
- if tidy == 0:
- shared_res[tidx] = 0.0
- cute.arch.sync_threads()
-
- # Get K tile count for reduction loop
- k_tile_cnt = gA_mkl.layout[3].shape
-
- # Main reduction loop over K tiles
- # Each thread in tidy processes a subset of K tiles
- for k_tile in range(tidy, k_tile_cnt, threads_per_k, unroll_full=True):
- # Load A tile and scale factors
- tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
- tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
+ for (int tile = 0; tile < num_tiles; tile++) {
+ int curr = tile & 1;
+ int next = (tile + 1) & 1;
+ int next_tile = tile + 1;
- # Load B tile and scale factors (B is broadcast across M)
- tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
- tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
+ // Prefetch next tile
+ if (next_tile < num_tiles) {
+ int k_base_next = next_tile * K_TILE_SIZE;
+ if (tid < 32) {
+ int k_off = k_base_next + tid * 16;
+ if (k_off < k_packed) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
+ ASYNC_COPY_16(dst, B + k_off);
+ }
+ }
+ if (tid < 4) {
+ int sf_off = k_base_next / 8 + tid * 16;
+ if (sf_off < k_sf) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
+ ASYNC_COPY_16(dst, SFB + sf_off);
+ }
+ }
+ ASYNC_COMMIT();
+ }
- # Create register tensors
- tArA = cute.make_rmem_tensor_like(tAgA, c_dtype)
- tBrB = cute.make_rmem_tensor_like(tBgB, c_dtype)
- tArSFA = cute.make_rmem_tensor_like(tAgSFA, accum_dtype)
- tBrSFB = cute.make_rmem_tensor_like(tBgSFB, accum_dtype)
+ // Wait for current tile
+ ASYNC_WAIT_ALL();
+ __syncthreads();
- # Load from global memory and convert types
- a_val = tAgA.load().to(c_dtype)
- b_val = tBgB.load().to(c_dtype)
- sfa_val = tAgSFA.load().to(accum_dtype)
- sfb_val = tBgSFB.load().to(accum_dtype)
+ if (valid) {
+ int k_base = tile * K_TILE_SIZE;
+ int k_local = tidx * 16;
+ int k_global = k_base + k_local;
+
+ if (k_global + 16 <= k_packed) {
+ // Load A directly (per-row, can't share)
+ uint4 a_vec = *reinterpret_cast<const uint4*>(A + row * k_packed + k_global);
+ // Load B from shared memory
+ uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
+
+ const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
+ const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
+
+ int sf_local = k_local / 8;
+ int sf_global = k_global / 8;
+
+ float sfa0 = decode_fp8(SFA[row * k_sf + sf_global]);
+ float sfa1 = decode_fp8(SFA[row * k_sf + sf_global + 1]);
+ float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
+ float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
+ float scale0 = sfa0 * sfb0;
+ float scale1 = sfa1 * sfb1;
+
+ __half2 acc0 = __float2half2_rn(0.0f);
+ __half2 acc1 = __float2half2_rn(0.0f);
+
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
+ }
+ sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
+
+ #pragma unroll
+ for (int i = 8; i < 16; i++) {
+ acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
+ }
+ sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
+ }
+ }
- # Store to register tensors
- tArA.store(a_val)
- tBrB.store(b_val)
- tArSFA.store(sfa_val)
- tBrSFB.store(sfb_val)
-
- # Compute block-scaled dot product for this K tile
- for i in cutlass.range_constexpr(blk_k):
- res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
+ __syncthreads();
+ }
- # Reduce across K dimension using atomic add to shared memory
- atomic_add_fp32(res[0], elem_pointer(shared_res, tidx))
- cute.arch.sync_threads()
+ if (!valid) return;
- # Final store to global memory (only thread 0 in K dimension)
- if tidy == 0:
- out = scalar_to_ssa(shared_res[tidx], cutlass.Float32)
- tCgC.store(out.to(cutlass.Float16))
+ // Warp shuffle reduction
+ #pragma unroll
+ for (int offset = 16; offset > 0; offset /= 2) {
+ sum += __shfl_down_sync(0xffffffff, sum, offset);
+ }
- return
+ if (tidx == 0) {
+ C[row] = __float2half(sum);
+ }
+ }
-
- @cute.jit
- def gemv_launcher(
- a_ptr: cute.Pointer,
- b_ptr: cute.Pointer,
- sfa_ptr: cute.Pointer,
- sfb_ptr: cute.Pointer,
- c_ptr: cute.Pointer,
- problem_size: tuple,
- ):
- """Host-side JIT function to prepare tensors and launch kernel."""
- m, _, k, l = problem_size
+ // Kernel for l > 1 with double-buffered B
+ __global__ __launch_bounds__(BLOCK_SIZE, 2)
+ void nvfp4_gemv_kernel_batched(
+ half* __restrict__ C,
+ const uint8_t* __restrict__ A,
+ const uint8_t* __restrict__ B,
+ const uint8_t* __restrict__ SFA,
+ const uint8_t* __restrict__ SFB,
+ int m, int k_packed, int k_sf, int l, int n_pad
+ ) {
+ __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
+ __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
- # Create A tensor: [m, k, l] K-major
- a_tensor = cute.make_tensor(
- a_ptr,
- cute.make_layout(
- (m, cute.assume(k, 32), l),
- stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
- ),
- )
+ const int tidx = threadIdx.x;
+ const int tidy = threadIdx.y;
+ const int tid = tidy * THREADS_PER_K + tidx;
+ const int row = blockIdx.x * THREADS_PER_M + tidy;
+ const int batch = blockIdx.z;
+ const bool valid = (row < m);
- # Create B tensor: [n_padded, k, l] K-major
- n_padded = 128
- b_tensor = cute.make_tensor(
- b_ptr,
- cute.make_layout(
- (n_padded, cute.assume(k, 32), l),
- stride=(cute.assume(k, 32), 1, cute.assume(n_padded * k, 32)),
- ),
- )
+ // Batch offsets
+ const size_t a_batch = batch * (size_t)(m * k_packed);
+ const size_t b_batch = batch * (size_t)(n_pad * k_packed);
+ const size_t sfa_batch = batch * (size_t)(m * k_sf);
+ const size_t sfb_batch = batch * (size_t)(n_pad * k_sf);
- # Create C tensor: [m, 1, l]
- c_tensor = cute.make_tensor(
- c_ptr,
- cute.make_layout(
- (cute.assume(m, 32), 1, l),
- stride=(1, 1, m)
- )
- )
+ const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
- # Create scale factor tensors with MMA layout
- sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
- sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
+ float sum = 0.0f;
- sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
- sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
+ // Load first tile
+ if (num_tiles > 0) {
+ if (tid < 32) {
+ int k_off = tid * 16;
+ if (k_off < k_packed) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
+ ASYNC_COPY_16(dst, B + b_batch + k_off);
+ }
+ }
+ if (tid < 4) {
+ int sf_off = tid * 16;
+ if (sf_off < k_sf) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
+ ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
+ }
+ }
+ ASYNC_COMMIT();
+ }
- # Compute grid dimensions
- grid = (
- cute.ceil_div(c_tensor.shape[0], threads_per_m),
- 1,
- c_tensor.shape[2],
- )
+ for (int tile = 0; tile < num_tiles; tile++) {
+ int curr = tile & 1;
+ int next = (tile + 1) & 1;
+ int next_tile = tile + 1;
+
+ // Prefetch next tile
+ if (next_tile < num_tiles) {
+ int k_base_next = next_tile * K_TILE_SIZE;
+ if (tid < 32) {
+ int k_off = k_base_next + tid * 16;
+ if (k_off < k_packed) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
+ ASYNC_COPY_16(dst, B + b_batch + k_off);
+ }
+ }
+ if (tid < 4) {
+ int sf_off = k_base_next / 8 + tid * 16;
+ if (sf_off < k_sf) {
+ uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
+ ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
+ }
+ }
+ ASYNC_COMMIT();
+ }
+
+ ASYNC_WAIT_ALL();
+ __syncthreads();
+
+ if (valid) {
+ int k_base = tile * K_TILE_SIZE;
+ int k_local = tidx * 16;
+ int k_global = k_base + k_local;
+
+ if (k_global + 16 <= k_packed) {
+ uint4 a_vec = *reinterpret_cast<const uint4*>(A + a_batch + row * k_packed + k_global);
+ uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
+
+ const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
+ const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
+
+ int sf_local = k_local / 8;
+ int sf_global = k_global / 8;
+
+ float sfa0 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global]);
+ float sfa1 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global + 1]);
+ float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
+ float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
+ float scale0 = sfa0 * sfb0;
+ float scale1 = sfa1 * sfb1;
+
+ __half2 acc0 = __float2half2_rn(0.0f);
+ __half2 acc1 = __float2half2_rn(0.0f);
+
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
+ }
+ sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
+
+ #pragma unroll
+ for (int i = 8; i < 16; i++) {
+ acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
+ }
+ sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
+ }
+ }
+
+ __syncthreads();
+ }
- # Launch kernel
- gemv_kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
- grid=grid,
- block=[threads_per_m, threads_per_k, 1],
- cluster=(1, 1, 1),
- )
+ if (!valid) return;
- return
+ #pragma unroll
+ for (int offset = 16; offset > 0; offset /= 2) {
+ sum += __shfl_down_sync(0xffffffff, sum, offset);
+ }
+
+ if (tidx == 0) {
+ C[batch * m + row] = __float2half(sum);
+ }
+ }
+ void run_nvfp4_gemv(
+ torch::Tensor C,
+ torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor SFA,
+ torch::Tensor SFB,
+ int m, int k, int l, int n_pad
+ ) {
+ int k_packed = k / 2;
+ int k_sf = k / SF_VEC_SIZE;
+
+ half* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
+ const uint8_t* a_ptr = A.data_ptr<uint8_t>();
+ const uint8_t* b_ptr = B.data_ptr<uint8_t>();
+ const uint8_t* sfa_ptr = SFA.data_ptr<uint8_t>();
+ const uint8_t* sfb_ptr = SFB.data_ptr<uint8_t>();
+
+ int blocks_m = (m + THREADS_PER_M - 1) / THREADS_PER_M;
+ dim3 block(THREADS_PER_K, THREADS_PER_M, 1);
- # Global cache for compiled kernel
- _compiled_kernel_cache = None
+ if (l == 1) {
+ dim3 grid(blocks_m, 1, 1);
+ nvfp4_gemv_kernel_single<<<grid, block>>>(
+ c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
+ m, k_packed, k_sf, n_pad
+ );
+ } else {
+ dim3 grid(blocks_m, 1, l);
+ nvfp4_gemv_kernel_batched<<<grid, block>>>(
+ c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
+ m, k_packed, k_sf, l, n_pad
+ );
+ }
+ }
+ #undef ASYNC_COPY_16
+ #undef ASYNC_COMMIT
+ #undef ASYNC_WAIT_ALL
+ '''
- def compile_kernel():
- """Compile the kernel once and cache it."""
- global _compiled_kernel_cache
-
- if _compiled_kernel_cache is not None:
- return _compiled_kernel_cache
-
- # Create placeholder pointers for compilation
- a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
- sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
-
- try:
- _compiled_kernel_cache = cute.compile(
- gemv_launcher, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)
+ CPP_SRC = r'''
+ #include <torch/extension.h>
+
+ void run_nvfp4_gemv(
+ torch::Tensor C,
+ torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor SFA,
+ torch::Tensor SFB,
+ int m, int k, int l, int n_pad
+ );
+ '''
+
+ _cuda_module = None
+
+ def get_cuda_module():
+ global _cuda_module
+ if _cuda_module is None:
+ _cuda_module = load_inline(
+ name='nvfp4_gemv_enhanced_v1',
+ cpp_sources=CPP_SRC,
+ cuda_sources=CUDA_SRC,
+ functions=['run_nvfp4_gemv'],
+ extra_cuda_cflags=[
+ '-O3',
+ '--use_fast_math',
+ '-std=c++17',
+ '-gencode=arch=compute_100a,code=sm_100a'
+ ],
+ verbose=False,
)
- except Exception as e:
- raise RuntimeError(f"Kernel compilation failed: {e}")
-
- return _compiled_kernel_cache
+ return _cuda_module
def custom_kernel(data: input_t) -> output_t:
- """
- Execute the block-scaled GEMV kernel.
+ a, b, sfa_ref, sfb_ref, _, _, c = data
- This implementation processes all batches in a single kernel launch.
+ module = get_cuda_module()
- Args:
- data: Tuple of (a, b, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c) tensors
- a: [m, k/2, l] - Input matrix in float4e2m1fn_x2
- b: [n_pad, k/2, l] - Input vector (padded to 128) in float4e2m1fn_x2
- sfa_ref: [m, sf_k, l] - Scale factors for A (not used)
- sfb_ref: [n_pad, sf_k, l] - Scale factors for B (not used)
- sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors for A (MMA layout)
- sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors for B (MMA layout)
- c: [m, 1, l] - Output vector in float16
+ m, k_packed, l = a.shape
+ k = k_packed * 2
+ n_pad = b.shape[0]
- Returns:
- Output tensor c with computed GEMV results
- """
- a, b, _, _, sfa_permuted, sfb_permuted, c = data
+ if not sfa_ref.is_cuda:
+ sfa_ref = sfa_ref.to(a.device)
+ if not sfb_ref.is_cuda:
+ sfb_ref = sfb_ref.to(a.device)
- # Compile kernel (uses cache if available)
- compiled_func = compile_kernel()
+ a_uint8 = a.view(torch.uint8)
+ b_uint8 = b.view(torch.uint8)
+ sfa_uint8 = sfa_ref.view(torch.uint8)
+ sfb_uint8 = sfb_ref.view(torch.uint8)
- # Get dimensions
- m, k_packed, l = a.shape
- k = k_packed * 2 # FP4 packed: 2 elements per byte
- n = 1 # GEMV
+ c_out = c.squeeze(1)
- # Create CuTe pointers from PyTorch tensors
- a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- sfa_ptr = make_ptr(sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
- sfb_ptr = make_ptr(sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
+ module.run_nvfp4_gemv(
+ c_out, a_uint8, b_uint8, sfa_uint8, sfb_uint8,
+ m, k, l, n_pad
+ )
- # Execute kernel - processes all batches in a single launch
- compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
-
return c
scrolls · 661 diff lines total

Best evidence level for this revision: reported

JSON