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
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-copy
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))fp4
1. Native FP4/FP8 conversion functionsfp8
__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 = uint4
uint4 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_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+ 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_moduledef 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