Skip to content
KernelIndex
Search⌘K

submission 71977

s.am._ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-71977?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
46.4µs
#258 of 678
2025-11-12

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:27fcff2d71f8fe5d2b79d204bc48c77125408c002f26a142d2ba0c60f84f08ac
license declaredunknown
license concludedunknown
authorss.am._
imported2026-08-15

Techniques

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

fp8const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;

Kernel source

sub_v2.py214 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# ---- C++ stub: declare the function so load_inline can bind it ----
gemv_cpp = r"""
#include <torch/extension.h>

// Forward declaration so PyTorch can bind it (definition is in the CUDA source).
torch::Tensor nvfp4_gemv_v1(torch::Tensor A,
                            torch::Tensor B,
                            torch::Tensor C,
                            torch::Tensor SFA,
                            torch::Tensor SFB);
"""

# ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----
gemv_cuda = r"""
#include <assert.h>
#include <cuda.h>
#include <stdio.h>
#include <cuda_runtime.h>

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>

#include <cuda_fp4.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>

// ---- gemv.h ----
struct Gemv_params {
    using index_t = uint64_t;

    int b, m, k, real_k;

    void *__restrict__ a_ptr;
    void *__restrict__ b_ptr;
    void *__restrict__ sfa_ptr;
    void *__restrict__ sfb_ptr;
    void *__restrict__ o_ptr;

    index_t a_batch_stride;
    index_t b_batch_stride;
    index_t sfa_batch_stride;
    index_t sfb_batch_stride;
    index_t o_batch_stride;

    index_t a_row_stride;
    index_t b_row_stride;
    index_t sfa_row_stride;
    index_t sfb_row_stride;
    index_t o_row_stride;
};

// ---- gemv_v1.cu ----
static constexpr int ROWS_PER_BLOCK  = 8;
static constexpr int THREADS_PER_ROW = 16;
static constexpr int BLOCK_SIZE      = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128

__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel(const __grid_constant__ Gemv_params params)
{
    const int tid      = threadIdx.x;
    const int rib      = tid / THREADS_PER_ROW;
    const int lane     = tid % THREADS_PER_ROW;
    const int batch    = blockIdx.z;
    const int row      = blockIdx.x * ROWS_PER_BLOCK + rib;

    const size_t A_batch_base   = static_cast<size_t>(batch) * params.a_batch_stride;
    const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
    const size_t B_batch_base   = static_cast<size_t>(batch) * params.b_batch_stride;
    const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
    const size_t C_batch_base   = static_cast<size_t>(batch) * params.o_batch_stride;

    const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;
    const __nv_fp8_e4m3*    rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;

    const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
    const __nv_fp8_e4m3*    vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;

    float sum = 0.f;

    for (int idx = 0; idx < params.k / THREADS_PER_ROW / 8; ++idx) {
        int base = idx * 16;
        const int base_id = (idx * THREADS_PER_ROW + lane) * 8;

        __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);
        __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);
        __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);
        __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);
        __half scale = __hmul(sfa, sfb);

        __half2 acc = __float2half2_rn(0.0f);

        #pragma unroll
        for (int i = 0; i < 8; ++i) {
            const int id = base_id + i;
            __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);
            __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);
            __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);
            __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);

            const __half2 a_h2 = __half2(a_raw);
            const __half2 b_h2 = __half2(b_raw);

            acc = __hfma2(a_h2, b_h2, acc);
        }

        __half fin = __hadd(__low2half(acc), __high2half(acc));
        __half h = __hmul(fin, scale);

        sum += __half2float(h);
    }

    unsigned mask = 0xffffffffu;
    sum += __shfl_down_sync(mask, sum, 8, 16);
    sum += __shfl_down_sync(mask, sum, 4, 16);
    sum += __shfl_down_sync(mask, sum, 2, 16);
    sum += __shfl_down_sync(mask, sum, 1, 16);

    if (lane == 0) {
        __half* out = (__half*)params.o_ptr + C_batch_base + row;
        out[0] = __float2half(sum);
    }
}

static inline void launch_kernel(Gemv_params &params, cudaStream_t stream)
{
    const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
    dim3 grid(grid_x, 1, params.b);
    dim3 block(BLOCK_SIZE, 1, 1);
    gemv_kernel<<<grid, block, 0, stream>>>(params);
}

// Python-facing function: sets up params and launches
torch::Tensor nvfp4_gemv_v1(torch::Tensor A,
                            torch::Tensor B,
                            torch::Tensor C,
                            torch::Tensor SFA,
                            torch::Tensor SFB)
{
    TORCH_CHECK(A.device().is_cuda(), "A must be CUDA");
    TORCH_CHECK(B.device().is_cuda(), "B must be CUDA");
    TORCH_CHECK(C.device().is_cuda(), "C must be CUDA");
    TORCH_CHECK(SFA.device().is_cuda(), "SFA must be CUDA");
    TORCH_CHECK(SFB.device().is_cuda(), "SFB must be CUDA");

    // Expect A: [M,K,L], B: [K,L], C: [M,1,L] or similar layout using provided strides.
    auto sizes = A.sizes();
    TORCH_CHECK(sizes.size() == 3, "A must be 3D [M,K,L]");
    const int64_t M = sizes[0];
    const int64_t K = sizes[1];
    const int64_t L = sizes[2];

    Gemv_params params{};
    params.b = static_cast<int>(L);
    params.m = static_cast<int>(M);
    params.k = static_cast<int>(K);
    params.real_k = static_cast<int>(K * 2);

    params.a_ptr  = A.data_ptr();
    params.b_ptr  = B.data_ptr();
    params.sfa_ptr= SFA.data_ptr();
    params.sfb_ptr= SFB.data_ptr();
    params.o_ptr  = C.data_ptr();

    params.a_batch_stride  = static_cast<uint64_t>(A.stride(2));
    params.b_batch_stride  = static_cast<uint64_t>(B.stride(2));
    params.sfa_batch_stride= static_cast<uint64_t>(SFA.stride(2));
    params.sfb_batch_stride= static_cast<uint64_t>(SFB.stride(2));
    params.o_batch_stride  = static_cast<uint64_t>(C.stride(2));

    params.a_row_stride  = static_cast<uint64_t>(A.stride(0));
    params.b_row_stride  = static_cast<uint64_t>(B.stride(0));
    params.sfa_row_stride= static_cast<uint64_t>(SFA.stride(0));
    params.sfb_row_stride= static_cast<uint64_t>(SFB.stride(0));
    params.o_row_stride  = static_cast<uint64_t>(C.stride(0));

    auto stream = at::cuda::getCurrentCUDAStream().stream();
    launch_kernel(params, stream);

    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "CUDA kernel failed: ", cudaGetErrorString(err));

    return C;
}
"""

# ---- build the module ----
nvfp4_module = load_inline(
    name="nvfp4_gemv",
    cpp_sources=[gemv_cpp],
    cuda_sources=[gemv_cuda],
    functions=["nvfp4_gemv_v1"],  # this exposes the function to Python
    extra_cuda_cflags=[
        "-std=c++17",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--ptxas-options=--gpu-name=sm_100a",
        "-O3",
        "-w",
        "--use_fast_math",
        "-allow-unsupported-compiler",
    ],
    extra_ldflags=["-lcuda", "-lcublas"],
    verbose=True,
)


def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
    return nvfp4_module.nvfp4_gemv_v1(a, b, c, sfa, sfb)
scrolls · 214 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 71063.

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
+ # ---- C++ stub: declare the function so load_inline can bind it ----
+ gemv_cpp = r"""
+ #include <torch/extension.h>
- # Kernel configuration parameters
- mma_tiler_mnk = (128, 1, 64) # Tile sizes for M, N, K dimensions
- 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
- sf_vec_size = 16 # Scale factor block size (16 elements share one scale)
- threads_per_cta = 128 # Number of threads per CUDA thread block
+ // Forward declaration so PyTorch can bind it (definition is in the CUDA source).
+ torch::Tensor nvfp4_gemv_v1(torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor C,
+ torch::Tensor SFA,
+ torch::Tensor SFB);
+ """
+ # ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----
+ gemv_cuda = r"""
+ #include <assert.h>
+ #include <cuda.h>
+ #include <stdio.h>
+ #include <cuda_runtime.h>
- # Helper function for ceiling division
- def ceil_div(a, b):
- return (a + b - 1) // b
+ #include <torch/extension.h>
+ #include <ATen/cuda/CUDAContext.h>
+ #include <c10/cuda/CUDAGuard.h>
+ #include <cuda_fp4.h>
+ #include <cuda_bf16.h>
+ #include <cuda_fp8.h>
- # The CuTe reference implementation for NVFP4 block-scaled GEMV
- @cute.kernel
- def kernel(
- mA_mkl: cute.Tensor,
- mB_nkl: cute.Tensor,
- mSFA_mkl: cute.Tensor,
- mSFB_nkl: cute.Tensor,
- mC_mnl: cute.Tensor,
- ):
- # Get CUDA block and thread indices
- bidx, bidy, bidz = cute.arch.block_idx()
- tidx, _, _ = cute.arch.thread_idx()
+ // ---- gemv.h ----
+ struct Gemv_params {
+ using index_t = uint64_t;
- # Extract the local tile for input matrix A (shape: [block_M, block_K, rest_M, rest_K, rest_L])
- gA_mkl = cute.local_tile(
- mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
- # Extract the local tile for scale factor tensor for A (same shape as gA_mkl)
- # Here, block_M = (32, 4); block_K = (16, 4)
- gSFA_mkl = cute.local_tile(
- mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
- # Extract the local tile for input matrix B (shape: [block_N, block_K, rest_N, rest_K, rest_L])
- gB_nkl = cute.local_tile(
- mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
- # Extract the local tile for scale factor tensor for B (same shape as gB_nkl)
- gSFB_nkl = cute.local_tile(
- mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
- # Extract the local tile for output matrix C (shape: [block_M, block_N, rest_M, rest_N, rest_L])
- gC_mnl = cute.local_tile(
- mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
- )
+ int b, m, k, real_k;
- # Select output element corresponding to this thread and block indices
- tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
- tCgC = cute.make_tensor(tCgC.iterator, 1)
- res = cute.zeros_like(tCgC, cutlass.Float32)
+ void *__restrict__ a_ptr;
+ void *__restrict__ b_ptr;
+ void *__restrict__ sfa_ptr;
+ void *__restrict__ sfb_ptr;
+ void *__restrict__ o_ptr;
- # Get the number of k tiles (depth dimension) for the reduction loop
- k_tile_cnt = gA_mkl.layout[3].shape
- for k_tile in range(k_tile_cnt):
- tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
- tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
- tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
- tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
+ index_t a_batch_stride;
+ index_t b_batch_stride;
+ index_t sfa_batch_stride;
+ index_t sfb_batch_stride;
+ index_t o_batch_stride;
- tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
- tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
- tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
- tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
+ index_t a_row_stride;
+ index_t b_row_stride;
+ index_t sfa_row_stride;
+ index_t sfb_row_stride;
+ index_t o_row_stride;
+ };
- # Load NVFP4 or FP8 values from global memory
- a_val_nvfp4 = tAgA.load()
- b_val_nvfp4 = tBgB.load()
- sfa_val_fp8 = tAgSFA.load()
- sfb_val_fp8 = tBgSFB.load()
+ // ---- gemv_v1.cu ----
+ static constexpr int ROWS_PER_BLOCK = 8;
+ static constexpr int THREADS_PER_ROW = 16;
+ static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128
- # Convert loaded values to float32 for computation (FFMA)
- a_val = a_val_nvfp4.to(cutlass.Float32)
- b_val = b_val_nvfp4.to(cutlass.Float32)
- sfa_val = sfa_val_fp8.to(cutlass.Float32)
- sfb_val = sfb_val_fp8.to(cutlass.Float32)
+ __global__ void __launch_bounds__(BLOCK_SIZE, 8)
+ gemv_kernel(const __grid_constant__ Gemv_params params)
+ {
+ const int tid = threadIdx.x;
+ const int rib = tid / THREADS_PER_ROW;
+ const int lane = tid % THREADS_PER_ROW;
+ const int batch = blockIdx.z;
+ const int row = blockIdx.x * ROWS_PER_BLOCK + rib;
- # Store the converted values to RMEM CuTe tensors
- tArA.store(a_val)
- tBrB.store(b_val)
- tArSFA.store(sfa_val)
- tBrSFB.store(sfb_val)
+ const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;
+ const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
+ const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;
+ const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
+ const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;
- # Iterate over SF vector tiles and compute the scale&matmul accumulation
- for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
- res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
+ const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;
+ const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
- # Store the final float16 result back to global memory
- tCgC.store(res.to(cutlass.Float16))
- return
+ const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
+ const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;
+ float sum = 0.f;
- @cute.jit
- def my_kernel(
- 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 GPU kernel.
- """
- m, _, k, l = problem_size
- # Create CuTe Tensor via pointer and problem size.
- 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)),
- ),
- )
- # We use n=128 to create the torch tensor to do fp4 computation via torch._scaled_mm
- # then copy torch tensor to cute tensor for cute customize kernel computation
- # therefore we need to ensure b_tensor has the right stride with this 128 padded size on n.
- n_padded_128 = 128
- b_tensor = cute.make_tensor(
- b_ptr,
- cute.make_layout(
- (n_padded_128, cute.assume(k, 32), l),
- stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
- ),
- )
- c_tensor = cute.make_tensor(
- c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
- )
- # Convert scale factor tensors to MMA layout
- # The layout matches Tensor Core requirements: (((32, 4), REST_M), ((SF_K, 4), REST_K), (1, REST_L))
- sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
- sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
+ for (int idx = 0; idx < params.k / THREADS_PER_ROW / 8; ++idx) {
+ int base = idx * 16;
+ const int base_id = (idx * THREADS_PER_ROW + lane) * 8;
- sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
- sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
+ __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);
+ __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);
+ __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);
+ __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);
+ __half scale = __hmul(sfa, sfb);
- # Compute grid dimensions
- # Grid is (M_blocks, 1, L) where:
- # - M_blocks = ceil(M / 128) to cover all output rows
- # - L = batch size
- grid = (
- cute.ceil_div(c_tensor.shape[0], 128),
- 1,
- c_tensor.shape[2],
- )
+ __half2 acc = __float2half2_rn(0.0f);
- # Launch the CUDA kernel
- kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
- grid=grid,
- block=[threads_per_cta, 1, 1],
- cluster=(1, 1, 1),
- )
- return
+ #pragma unroll
+ for (int i = 0; i < 8; ++i) {
+ const int id = base_id + i;
+ __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);
+ __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);
+ __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);
+ __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);
+ const __half2 a_h2 = __half2(a_raw);
+ const __half2 b_h2 = __half2(b_raw);
- # Global cache for compiled kernel
- _compiled_kernel_cache = None
+ acc = __hfma2(a_h2, b_h2, acc);
+ }
+ __half fin = __hadd(__low2half(acc), __high2half(acc));
+ __half h = __hmul(fin, scale);
- # This function is used to compile the kernel once and cache it and then allow users to
- # run the kernel multiple times to get more accurate timing results.
- def compile_kernel():
- """
- Compile the kernel once and cache it.
- This should be called before any timing measurements.
+ sum += __half2float(h);
+ }
- Returns:
- The compiled kernel function
- """
- global _compiled_kernel_cache
+ unsigned mask = 0xffffffffu;
+ sum += __shfl_down_sync(mask, sum, 8, 16);
+ sum += __shfl_down_sync(mask, sum, 4, 16);
+ sum += __shfl_down_sync(mask, sum, 2, 16);
+ sum += __shfl_down_sync(mask, sum, 1, 16);
- if _compiled_kernel_cache is not None:
- return _compiled_kernel_cache
+ if (lane == 0) {
+ __half* out = (__half*)params.o_ptr + C_batch_base + row;
+ out[0] = __float2half(sum);
+ }
+ }
- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
- 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)
+ static inline void launch_kernel(Gemv_params &params, cudaStream_t stream)
+ {
+ const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
+ dim3 grid(grid_x, 1, params.b);
+ dim3 block(BLOCK_SIZE, 1, 1);
+ gemv_kernel<<<grid, block, 0, stream>>>(params);
+ }
- # Compile the kernel
- _compiled_kernel_cache = cute.compile(
- my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)
- )
+ // Python-facing function: sets up params and launches
+ torch::Tensor nvfp4_gemv_v1(torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor C,
+ torch::Tensor SFA,
+ torch::Tensor SFB)
+ {
+ TORCH_CHECK(A.device().is_cuda(), "A must be CUDA");
+ TORCH_CHECK(B.device().is_cuda(), "B must be CUDA");
+ TORCH_CHECK(C.device().is_cuda(), "C must be CUDA");
+ TORCH_CHECK(SFA.device().is_cuda(), "SFA must be CUDA");
+ TORCH_CHECK(SFB.device().is_cuda(), "SFB must be CUDA");
- return _compiled_kernel_cache
+ // Expect A: [M,K,L], B: [K,L], C: [M,1,L] or similar layout using provided strides.
+ auto sizes = A.sizes();
+ TORCH_CHECK(sizes.size() == 3, "A must be 3D [M,K,L]");
+ const int64_t M = sizes[0];
+ const int64_t K = sizes[1];
+ const int64_t L = sizes[2];
+ Gemv_params params{};
+ params.b = static_cast<int>(L);
+ params.m = static_cast<int>(M);
+ params.k = static_cast<int>(K);
+ params.real_k = static_cast<int>(K * 2);
- def custom_kernel(data: input_t) -> output_t:
- """
- Execute the block-scaled GEMV kernel.
+ params.a_ptr = A.data_ptr();
+ params.b_ptr = B.data_ptr();
+ params.sfa_ptr= SFA.data_ptr();
+ params.sfb_ptr= SFB.data_ptr();
+ params.o_ptr = C.data_ptr();
- This is the main entry point called by the evaluation framework.
- It converts PyTorch tensors to CuTe tensors, launches the kernel,
- and returns the result.
+ params.a_batch_stride = static_cast<uint64_t>(A.stride(2));
+ params.b_batch_stride = static_cast<uint64_t>(B.stride(2));
+ params.sfa_batch_stride= static_cast<uint64_t>(SFA.stride(2));
+ params.sfb_batch_stride= static_cast<uint64_t>(SFB.stride(2));
+ params.o_batch_stride = static_cast<uint64_t>(C.stride(2));
- Args:
- data: Tuple of (a, b, sfa_cpu, sfb_cpu, c) PyTorch tensors
- a: [m, k, l] - Input matrix in float4e2m1fn
- b: [1, k, l] - Input vector in float4e2m1fn
- sfa_cpu: [m, k, l] - Scale factors in float8_e4m3fn
- sfb_cpu: [1, k, l] - Scale factors in float8_e4m3fn
- sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn
- sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
- c: [m, 1, l] - Output vector in float16
+ params.a_row_stride = static_cast<uint64_t>(A.stride(0));
+ params.b_row_stride = static_cast<uint64_t>(B.stride(0));
+ params.sfa_row_stride= static_cast<uint64_t>(SFA.stride(0));
+ params.sfb_row_stride= static_cast<uint64_t>(SFB.stride(0));
+ params.o_row_stride = static_cast<uint64_t>(C.stride(0));
- Returns:
- Output tensor c with computed GEMV results
- """
- a, b, _, _, sfa_permuted, sfb_permuted, c = data
+ auto stream = at::cuda::getCurrentCUDAStream().stream();
+ launch_kernel(params, stream);
- # Ensure kernel is compiled (will use cached version if available)
- # To avoid the compilation overhead, we compile the kernel once and cache it.
- compiled_func = compile_kernel()
+ cudaError_t err = cudaGetLastError();
+ TORCH_CHECK(err == cudaSuccess, "CUDA kernel failed: ", cudaGetErrorString(err));
- # Get dimensions from MxKxL layout
- m, k, l = a.shape
- # Torch use e2m1_x2 data type, thus k is halved
- k = k * 2
- # GEMV N dimension is always 1
- n = 1
+ return C;
+ }
+ """
- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
- 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
- )
+ # ---- build the module ----
+ nvfp4_module = load_inline(
+ name="nvfp4_gemv",
+ cpp_sources=[gemv_cpp],
+ cuda_sources=[gemv_cuda],
+ functions=["nvfp4_gemv_v1"], # this exposes the function to Python
+ extra_cuda_cflags=[
+ "-std=c++17",
+ "-gencode=arch=compute_100a,code=sm_100a",
+ "--ptxas-options=--gpu-name=sm_100a",
+ "-O3",
+ "-w",
+ "--use_fast_math",
+ "-allow-unsupported-compiler",
+ ],
+ extra_ldflags=["-lcuda", "-lcublas"],
+ verbose=True,
+ )
- # Execute the compiled kernel
- compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
- return c
+ def custom_kernel(data: input_t) -> output_t:
+ a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
+ return nvfp4_module.nvfp4_gemv_v1(a, b, c, sfa, sfb)
scrolls · 419 diff lines total

Best evidence level for this revision: reported

JSON