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
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.
fp8
const __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 ¶ms, 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_inlinefrom 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 ¶ms, 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