submission 75888
mdouglas · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 416 lines, June 9 Researcher Reciprocity License v1.0.
submission_cuda.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-75888?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:9db8cb51bc93b7e5f7feae47e94a8c0a95246da6be54c8197a1ecab3099b3553
license declaredunknown
license concludedunknown
authorsmdouglas
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const uint8_t* __restrict__ a, // [M, K//2] packed FP4 (2 per byte)fp8
const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)shared-memory
extern __shared__ uint8_t smem[];vector-width = uint4
reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];Kernel source
submission_cuda.py416 lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# CUDA kernel code for NVFP4 block-scaled GEMV
# Uses native Blackwell (sm_100a) hardware intrinsics for FP4/FP8 conversion
cuda_source = """
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <ATen/cuda/Exceptions.h>
#include <ATen/cuda/CUDAContext.h>
// NVFP4 is 4-bit float (e2m1): 1 sign bit, 2 exponent bits, 1 mantissa bit
// Stored as 2 values per byte
// Scale factors are FP8 (e4m3) for every 16 FP4 values
// Warp reduction using shuffle - optimized for Blackwell
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset /= 2) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
return val;
}
// Batched kernel: processes all L batches in one launch for better efficiency.
// Each block handles one M value, each warp within handles one L batch.
// Inputs have native PyTorch strides from .permute() - K dimension has stride 1.
__global__ void nvfp4_gemv_batched_kernel(
const uint8_t* __restrict__ a, // [M, K//2, L] with strides (K_half, 1, M*K_half)
const uint8_t* __restrict__ b, // [N, K//2, L] with strides (K_half, 1, N*K_half), N=128 padded
const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)
const __nv_fp8_e4m3* __restrict__ sfb, // [N, K//16, L] with strides (K_div_16, 1, N*K_div_16)
half* __restrict__ c, // [M, 1, L] output FP16
int M,
int K,
int L
) {
// Shared memory layout: B vectors [L, K/2], sfb [L, K/16]
extern __shared__ uint8_t smem[];
int K_half = K / 2;
int K_div_16 = K / 16;
uint8_t* sb = smem; // B vectors: L × K/2 bytes
__nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + L * K_half); // Scale factors B: L × K/16
int tid = threadIdx.x;
int warp_id = threadIdx.x / 32;
int lane = threadIdx.x % 32;
int m = blockIdx.x; // Each block handles one M value
if (m >= M) return;
const int N_padded = 128; // B is padded to 128 rows for torch._scaled_mm
// Cooperatively load all L B vectors into shared memory
// B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half (strides: K_half, 1, N_padded*K_half)
// We only need n=0 (the actual vector, rest is padding)
for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
// b[0, k, l] at offset: 0*K_half + k + l*N_padded*K_half
sb[l * K_half + k] = b[k + l * N_padded * K_half];
}
// Cooperatively load scale factors for B
for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
// sfb[0, k, l] at offset: 0*K_div_16 + k + l*N_padded*K_div_16
ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
}
__syncthreads();
// Each block processes multiple M rows, each warp processes multiple (m, l) pairs
const int M_ROWS_PER_BLOCK = 4;
const int WARPS_PER_BLOCK = blockDim.x / 32;
int m_start = blockIdx.x * M_ROWS_PER_BLOCK;
// Each warp processes all L batches for M_ROWS_PER_BLOCK / WARPS_PER_BLOCK M rows
int m_rows_per_warp = (M_ROWS_PER_BLOCK + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK;
for (int m_idx = 0; m_idx < m_rows_per_warp; m_idx++) {
int m = m_start + warp_id + m_idx * WARPS_PER_BLOCK;
if (m >= M) break;
// Process all L batches for this M row
for (int l = 0; l < L; l++) {
float sum = 0.0f;
// K dimension has stride 1, enabling coalesced access.
const int a_base = m * K_half + l * M * K_half;
const int sfa_base = m * K_div_16 + l * M * K_div_16;
const uint8_t* sb_row = &sb[l * K_half];
const __nv_fp8_e4m3* ssfb_row = &ssfb[l * K_div_16];
// Process in blocks of 8 bytes - each scale factor covers 8 bytes (16 FP4 values).
for (int scale_block = lane; scale_block < K_div_16; scale_block += 32) {
half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block].__x, __NV_E4M3).x);
half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block].__x, __NV_E4M3).x);
half combined_scale = scale_a * scale_b;
__half2 scale2 = __half2half2(combined_scale);
// Load all 8 bytes at once using uint2
int k_byte_base = scale_block * 8;
const uint2 a_data = *reinterpret_cast<const uint2*>(&a[a_base + k_byte_base]);
const uint2 b_data = *reinterpret_cast<const uint2*>(&sb_row[k_byte_base]);
const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);
// Process all 8 bytes with the same scale
#pragma unroll
for (int i = 0; i < 8; i++) {
__half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
__half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);
__half2 products = __hmul2(a_vals, b_vals);
__half2 scaled = __hmul2(products, scale2);
sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
}
}
sum = warp_reduce_sum(sum);
// c_ref has shape [M, 1, L] with strides (1, 1, M) from permute
// So c_ref[m, 0, l] is at linear offset: m + l*M
if (lane == 0) {
c[m + l * M] = __float2half(sum);
}
}
}
}
__global__ void nvfp4_gemv_kernel(
const uint8_t* __restrict__ a, // [M, K//2] packed FP4 (2 per byte)
const uint8_t* __restrict__ b, // [1, K//2] packed FP4
const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16] FP8 scale factors for A
const __nv_fp8_e4m3* __restrict__ sfb, // [1, K//16] FP8 scale factors for B
half* __restrict__ c, // [M, 1] output FP16
int M,
int K
) {
// 8 warps per M row for better memory latency hiding
// 2 M rows per block
const int WARPS_PER_M_ROW = 8;
const int WARPS_PER_BLOCK = blockDim.x / 32;
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2
// Shared memory: B vector, scale factors, and partial sums for reduction
extern __shared__ uint8_t smem[];
uint8_t* sb = smem;
__nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
float* partial_sums = reinterpret_cast<float*>(ssfb + K/16);
int tid = threadIdx.x;
int warp_id = tid / 32;
int lane = tid % 32;
int K_half = K / 2;
int K_div_16 = K / 16;
// Cooperatively load B vector into shared memory (vectorized)
int num_vec_loads = K_half / 16;
for (int i = tid; i < num_vec_loads; i += blockDim.x) {
reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];
}
int vec_bytes = num_vec_loads * 16;
for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {
sb[i] = b[i];
}
// Cooperatively load scale factors for B
for (int i = tid; i < K_div_16; i += blockDim.x) {
ssfb[i] = sfb[i];
}
__syncthreads();
// Each block processes M_ROWS_PER_BLOCK M rows
int m_base = blockIdx.x * M_ROWS_PER_BLOCK;
// Which M row within the block does this warp contribute to?
int m_local = warp_id / WARPS_PER_M_ROW;
int m = m_base + m_local;
if (m >= M) return;
// Which K chunk does this warp handle?
int warp_in_m_group = warp_id % WARPS_PER_M_ROW;
// Process by scale blocks instead of bytes - each scale covers 8 bytes (16 FP4 values)
int scales_per_warp = K_div_16 / WARPS_PER_M_ROW; // 1024 / 8 = 128 scales per warp
int scale_start = warp_in_m_group * scales_per_warp;
int scale_end = scale_start + scales_per_warp;
float sum = 0.0f;
// Loop over scale blocks - each iteration processes 8 bytes covered by one scale
// With 8 warps: 128 scales / 32 threads = 4 iterations per thread (down from 32!)
for (int scale_block = scale_start + lane; scale_block < scale_end; scale_block += 32) {
// Load scale factors ONCE for this block of 8 bytes
half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block].__x, __NV_E4M3).x);
half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block].__x, __NV_E4M3).x);
half combined_scale = scale_a * scale_b;
__half2 scale2 = __half2half2(combined_scale);
// Process all 8 bytes (8 fp4x2 pairs) covered by this scale factor
int k_byte_base = scale_block * 8;
// Load 8 bytes at once using uint2, then reinterpret as fp4x2 array
const uint2 a_data = *reinterpret_cast<const uint2*>(&a[m * K_half + k_byte_base]);
const uint2 b_data = *reinterpret_cast<const uint2*>(&sb[k_byte_base]);
const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);
const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);
// Process all 8 bytes with the same scale
#pragma unroll
for (int i = 0; i < 8; i++) {
__half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);
__half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);
__half2 products = __hmul2(a_vals, b_vals);
__half2 scaled = __hmul2(products, scale2);
sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);
sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);
}
}
// Intra-warp reduction
sum = warp_reduce_sum(sum);
// Store partial sum to shared memory
if (lane == 0) {
partial_sums[warp_id] = sum;
}
__syncthreads();
// Final reduction: first warp of each M group reduces the partial sums
if (warp_in_m_group == 0 && lane < WARPS_PER_M_ROW) {
float final_sum = partial_sums[m_local * WARPS_PER_M_ROW + lane];
// Reduce across the 8 partial sums
#pragma unroll
for (int offset = 4; offset > 0; offset /= 2) {
final_sum += __shfl_down_sync(0xffffffff, final_sum, offset);
}
if (lane == 0) {
c[m] = __float2half(final_sum);
}
}
}
void nvfp4_gemv_cuda(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M,
int K
) {
// 8 warps per M row, 2 M rows per block
const int WARPS_PER_M_ROW = 8;
const int WARPS_PER_BLOCK = 16;
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2
const int threads = WARPS_PER_BLOCK * 32; // 512 threads
const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
// Shared memory: B vector + sfb + partial sums
const int smem_size = K / 2 + K / 16 + WARPS_PER_BLOCK * sizeof(float);
// Get current CUDA stream from PyTorch
cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());
nvfp4_gemv_kernel<<<blocks, threads, smem_size, stream>>>(
a.data_ptr<uint8_t>(),
b.data_ptr<uint8_t>(),
reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<half*>(c.data_ptr<at::Half>()),
M, K
);
// Check for kernel launch errors
AT_CUDA_CHECK(cudaGetLastError());
}
void nvfp4_gemv_batched_cuda(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M,
int K,
int L
) {
// Each block handles 4 M rows, 8 warps process all (m, l) pairs
const int M_ROWS_PER_BLOCK = 4;
const int threads = 256;
const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
// Shared memory: B vectors (L × K/2) and sfb (L × K/16)
const int K_half = K / 2;
const int K_div_16 = K / 16;
const int smem_size = L * K_half + L * K_div_16;
// Get current CUDA stream from PyTorch
cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());
nvfp4_gemv_batched_kernel<<<blocks, threads, smem_size, stream>>>(
a.data_ptr<uint8_t>(),
b.data_ptr<uint8_t>(),
reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<half*>(c.data_ptr<at::Half>()),
M, K, L
);
// Check for kernel launch errors
AT_CUDA_CHECK(cudaGetLastError());
}
"""
cpp_source = """
void nvfp4_gemv_cuda(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M,
int K
);
void nvfp4_gemv_batched_cuda(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M,
int K,
int L
);
"""
# Compile the CUDA extension inline
nvfp4_gemv_module = load_inline(
name='nvfp4_gemv',
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
functions=['nvfp4_gemv_cuda', 'nvfp4_gemv_batched_cuda'],
verbose=True,
extra_cuda_cflags=[
'-O3',
'--use_fast_math',
'-arch=sm_100a',
'--std=c++17',
'-U__CUDA_NO_HALF_OPERATORS__', # Enable half operators
'-U__CUDA_NO_HALF_CONVERSIONS__', # Enable half conversions
],
)
def custom_kernel(data: input_t) -> output_t:
"""
Custom CUDA implementation of NVFP4 block-scaled GEMV.
Uses separate kernels optimized for L=1 and L>1 cases.
"""
import torch
a_ref, b_ref, sfa, sfb, _, _, c_ref = data
M, K_half, L = a_ref.shape
K = K_half * 2 # Each byte contains 2 FP4 values
if L == 1:
# Use single-batch optimized kernel.
# Pass c_ref directly - kernel writes c[m] which maps to c_ref[m, 0, 0].
a_bytes = a_ref[:, :, 0].view(torch.uint8).contiguous()
b_bytes = b_ref[0, :, 0].view(torch.uint8).contiguous()
nvfp4_gemv_module.nvfp4_gemv_cuda(
a_bytes,
b_bytes,
sfa[:, :, 0].contiguous(),
sfb[0, :, 0].contiguous(),
c_ref,
M, K
)
else:
# Use batched kernel for L>1 - processes all batches in one launch.
# Original layout already has K contiguous (stride 1) from creation
# as (L, M, K).permute(1, 2, 0), so no additional permute is needed.
a_bytes = a_ref.view(torch.uint8)
b_bytes = b_ref.view(torch.uint8)
nvfp4_gemv_module.nvfp4_gemv_batched_cuda(
a_bytes,
b_bytes,
sfa,
sfb,
c_ref,
M, K, L
)
return c_ref
scrolls · 416 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 71049.
- import torch+ from torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t- torch._dynamo.config.cache_size_limit = 32+ # CUDA kernel code for NVFP4 block-scaled GEMV+ # Uses native Blackwell (sm_100a) hardware intrinsics for FP4/FP8 conversion+ cuda_source = """+ #include <torch/extension.h>+ #include <cuda_runtime.h>+ #include <cuda_fp16.h>+ #include <cuda_fp4.h>+ #include <cuda_fp8.h>+ #include <ATen/cuda/Exceptions.h>+ #include <ATen/cuda/CUDAContext.h>- # Convert all scale factors to blocked formats+ // NVFP4 is 4-bit float (e2m1): 1 sign bit, 2 exponent bits, 1 mantissa bit+ // Stored as 2 values per byte+ // Scale factors are FP8 (e4m3) for every 16 FP4 values- @torch.compile(dynamic=False, fullgraph=True)- def to_blocked_3d(input_matrix):- # input_matrix is rows x cols x l- rows, cols, l = input_matrix.shape+ // Warp reduction using shuffle - optimized for Blackwell+ __device__ __forceinline__ float warp_reduce_sum(float val) {+ #pragma unroll+ for (int offset = 16; offset > 0; offset /= 2) {+ val += __shfl_down_sync(0xffffffff, val, offset);+ }+ return val;+ }- data = input_matrix.permute(2, 0, 1)+ // Batched kernel: processes all L batches in one launch for better efficiency.+ // Each block handles one M value, each warp within handles one L batch.+ // Inputs have native PyTorch strides from .permute() - K dimension has stride 1.+ __global__ void nvfp4_gemv_batched_kernel(+ const uint8_t* __restrict__ a, // [M, K//2, L] with strides (K_half, 1, M*K_half)+ const uint8_t* __restrict__ b, // [N, K//2, L] with strides (K_half, 1, N*K_half), N=128 padded+ const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16, L] with strides (K_div_16, 1, M*K_div_16)+ const __nv_fp8_e4m3* __restrict__ sfb, // [N, K//16, L] with strides (K_div_16, 1, N*K_div_16)+ half* __restrict__ c, // [M, 1, L] output FP16+ int M,+ int K,+ int L+ ) {+ // Shared memory layout: B vectors [L, K/2], sfb [L, K/16]+ extern __shared__ uint8_t smem[];+ int K_half = K / 2;+ int K_div_16 = K / 16;- return data.view(l, rows // 128, 128, cols // 4, 4) \- .transpose(2, 3) \- .reshape(l, -1, 4, 32, 4) \- .transpose(2, 3) \- .flatten(1)+ uint8_t* sb = smem; // B vectors: L × K/2 bytes+ __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + L * K_half); // Scale factors B: L × K/16- @torch.compile(dynamic=False, mode="max-autotune-no-cudagraphs")- def batched_gemv_impl(a_ref, b_ref, c_ref, sfa, sfb):- _, _, l = b_ref.shape+ int tid = threadIdx.x;+ int warp_id = threadIdx.x / 32;+ int lane = threadIdx.x % 32;+ int m = blockIdx.x; // Each block handles one M value- sfa_blocked = to_blocked_3d(sfa)- sfb_blocked = to_blocked_3d(sfb)+ if (m >= M) return;- for l_idx in range(l):- c_ref[:, 0, l_idx] = torch._scaled_mm(- a_ref[..., l_idx],- b_ref[..., l_idx].t(),- sfa_blocked[l_idx, ...],- sfb_blocked[l_idx, ...],- bias=None,- out_dtype=torch.float16,- )[:, 0]+ const int N_padded = 128; // B is padded to 128 rows for torch._scaled_mm- return c_ref+ // Cooperatively load all L B vectors into shared memory+ // B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half (strides: K_half, 1, N_padded*K_half)+ // We only need n=0 (the actual vector, rest is padding)+ for (int kl = tid; kl < K_half * L; kl += blockDim.x) {+ int k = kl / L;+ int l = kl % L;+ // b[0, k, l] at offset: 0*K_half + k + l*N_padded*K_half+ sb[l * K_half + k] = b[k + l * N_padded * K_half];+ }- def custom_kernel(- data: input_t,- ) -> output_t:+ // Cooperatively load scale factors for B+ for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {+ int k = kl / L;+ int l = kl % L;+ // sfb[0, k, l] at offset: 0*K_div_16 + k + l*N_padded*K_div_16+ ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];+ }++ __syncthreads();++ // Each block processes multiple M rows, each warp processes multiple (m, l) pairs+ const int M_ROWS_PER_BLOCK = 4;+ const int WARPS_PER_BLOCK = blockDim.x / 32;+ int m_start = blockIdx.x * M_ROWS_PER_BLOCK;++ // Each warp processes all L batches for M_ROWS_PER_BLOCK / WARPS_PER_BLOCK M rows+ int m_rows_per_warp = (M_ROWS_PER_BLOCK + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK;++ for (int m_idx = 0; m_idx < m_rows_per_warp; m_idx++) {+ int m = m_start + warp_id + m_idx * WARPS_PER_BLOCK;+ if (m >= M) break;++ // Process all L batches for this M row+ for (int l = 0; l < L; l++) {+ float sum = 0.0f;++ // K dimension has stride 1, enabling coalesced access.+ const int a_base = m * K_half + l * M * K_half;+ const int sfa_base = m * K_div_16 + l * M * K_div_16;+ const uint8_t* sb_row = &sb[l * K_half];+ const __nv_fp8_e4m3* ssfb_row = &ssfb[l * K_div_16];++ // Process in blocks of 8 bytes - each scale factor covers 8 bytes (16 FP4 values).+ for (int scale_block = lane; scale_block < K_div_16; scale_block += 32) {+ half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block].__x, __NV_E4M3).x);+ half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block].__x, __NV_E4M3).x);+ half combined_scale = scale_a * scale_b;+ __half2 scale2 = __half2half2(combined_scale);++ // Load all 8 bytes at once using uint2+ int k_byte_base = scale_block * 8;+ const uint2 a_data = *reinterpret_cast<const uint2*>(&a[a_base + k_byte_base]);+ const uint2 b_data = *reinterpret_cast<const uint2*>(&sb_row[k_byte_base]);++ const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);+ const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);++ // Process all 8 bytes with the same scale+ #pragma unroll+ for (int i = 0; i < 8; i++) {+ __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);+ __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);++ __half2 products = __hmul2(a_vals, b_vals);+ __half2 scaled = __hmul2(products, scale2);++ sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);+ sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);+ }+ }++ sum = warp_reduce_sum(sum);++ // c_ref has shape [M, 1, L] with strides (1, 1, M) from permute+ // So c_ref[m, 0, l] is at linear offset: m + l*M+ if (lane == 0) {+ c[m + l * M] = __float2half(sum);+ }+ }+ }+ }++ __global__ void nvfp4_gemv_kernel(+ const uint8_t* __restrict__ a, // [M, K//2] packed FP4 (2 per byte)+ const uint8_t* __restrict__ b, // [1, K//2] packed FP4+ const __nv_fp8_e4m3* __restrict__ sfa, // [M, K//16] FP8 scale factors for A+ const __nv_fp8_e4m3* __restrict__ sfb, // [1, K//16] FP8 scale factors for B+ half* __restrict__ c, // [M, 1] output FP16+ int M,+ int K+ ) {+ // 8 warps per M row for better memory latency hiding+ // 2 M rows per block+ const int WARPS_PER_M_ROW = 8;+ const int WARPS_PER_BLOCK = blockDim.x / 32;+ const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2++ // Shared memory: B vector, scale factors, and partial sums for reduction+ extern __shared__ uint8_t smem[];+ uint8_t* sb = smem;+ __nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);+ float* partial_sums = reinterpret_cast<float*>(ssfb + K/16);++ int tid = threadIdx.x;+ int warp_id = tid / 32;+ int lane = tid % 32;++ int K_half = K / 2;+ int K_div_16 = K / 16;++ // Cooperatively load B vector into shared memory (vectorized)+ int num_vec_loads = K_half / 16;+ for (int i = tid; i < num_vec_loads; i += blockDim.x) {+ reinterpret_cast<uint4*>(sb)[i] = reinterpret_cast<const uint4*>(b)[i];+ }+ int vec_bytes = num_vec_loads * 16;+ for (int i = tid + vec_bytes; i < K_half; i += blockDim.x) {+ sb[i] = b[i];+ }++ // Cooperatively load scale factors for B+ for (int i = tid; i < K_div_16; i += blockDim.x) {+ ssfb[i] = sfb[i];+ }++ __syncthreads();++ // Each block processes M_ROWS_PER_BLOCK M rows+ int m_base = blockIdx.x * M_ROWS_PER_BLOCK;++ // Which M row within the block does this warp contribute to?+ int m_local = warp_id / WARPS_PER_M_ROW;+ int m = m_base + m_local;++ if (m >= M) return;++ // Which K chunk does this warp handle?+ int warp_in_m_group = warp_id % WARPS_PER_M_ROW;++ // Process by scale blocks instead of bytes - each scale covers 8 bytes (16 FP4 values)+ int scales_per_warp = K_div_16 / WARPS_PER_M_ROW; // 1024 / 8 = 128 scales per warp+ int scale_start = warp_in_m_group * scales_per_warp;+ int scale_end = scale_start + scales_per_warp;++ float sum = 0.0f;++ // Loop over scale blocks - each iteration processes 8 bytes covered by one scale+ // With 8 warps: 128 scales / 32 threads = 4 iterations per thread (down from 32!)+ for (int scale_block = scale_start + lane; scale_block < scale_end; scale_block += 32) {+ // Load scale factors ONCE for this block of 8 bytes+ half scale_a = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block].__x, __NV_E4M3).x);+ half scale_b = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block].__x, __NV_E4M3).x);+ half combined_scale = scale_a * scale_b;+ __half2 scale2 = __half2half2(combined_scale);++ // Process all 8 bytes (8 fp4x2 pairs) covered by this scale factor+ int k_byte_base = scale_block * 8;++ // Load 8 bytes at once using uint2, then reinterpret as fp4x2 array+ const uint2 a_data = *reinterpret_cast<const uint2*>(&a[m * K_half + k_byte_base]);+ const uint2 b_data = *reinterpret_cast<const uint2*>(&sb[k_byte_base]);++ const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_data);+ const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_data);++ // Process all 8 bytes with the same scale+ #pragma unroll+ for (int i = 0; i < 8; i++) {+ __half2 a_vals = __nv_cvt_fp4x2_to_halfraw2(a_fp4x2[i], __NV_E2M1);+ __half2 b_vals = __nv_cvt_fp4x2_to_halfraw2(b_fp4x2[i], __NV_E2M1);++ __half2 products = __hmul2(a_vals, b_vals);+ __half2 scaled = __hmul2(products, scale2);++ sum = __fmaf_rn(__half2float(scaled.x), 1.0f, sum);+ sum = __fmaf_rn(__half2float(scaled.y), 1.0f, sum);+ }+ }++ // Intra-warp reduction+ sum = warp_reduce_sum(sum);++ // Store partial sum to shared memory+ if (lane == 0) {+ partial_sums[warp_id] = sum;+ }++ __syncthreads();++ // Final reduction: first warp of each M group reduces the partial sums+ if (warp_in_m_group == 0 && lane < WARPS_PER_M_ROW) {+ float final_sum = partial_sums[m_local * WARPS_PER_M_ROW + lane];+ // Reduce across the 8 partial sums+ #pragma unroll+ for (int offset = 4; offset > 0; offset /= 2) {+ final_sum += __shfl_down_sync(0xffffffff, final_sum, offset);+ }++ if (lane == 0) {+ c[m] = __float2half(final_sum);+ }+ }+ }++ void nvfp4_gemv_cuda(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa,+ torch::Tensor sfb,+ torch::Tensor c,+ int M,+ int K+ ) {+ // 8 warps per M row, 2 M rows per block+ const int WARPS_PER_M_ROW = 8;+ const int WARPS_PER_BLOCK = 16;+ const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW; // = 2+ const int threads = WARPS_PER_BLOCK * 32; // 512 threads+ const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;++ // Shared memory: B vector + sfb + partial sums+ const int smem_size = K / 2 + K / 16 + WARPS_PER_BLOCK * sizeof(float);++ // Get current CUDA stream from PyTorch+ cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());++ nvfp4_gemv_kernel<<<blocks, threads, smem_size, stream>>>(+ a.data_ptr<uint8_t>(),+ b.data_ptr<uint8_t>(),+ reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),+ reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),+ reinterpret_cast<half*>(c.data_ptr<at::Half>()),+ M, K+ );++ // Check for kernel launch errors+ AT_CUDA_CHECK(cudaGetLastError());+ }++ void nvfp4_gemv_batched_cuda(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa,+ torch::Tensor sfb,+ torch::Tensor c,+ int M,+ int K,+ int L+ ) {+ // Each block handles 4 M rows, 8 warps process all (m, l) pairs+ const int M_ROWS_PER_BLOCK = 4;+ const int threads = 256;+ const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;++ // Shared memory: B vectors (L × K/2) and sfb (L × K/16)+ const int K_half = K / 2;+ const int K_div_16 = K / 16;+ const int smem_size = L * K_half + L * K_div_16;++ // Get current CUDA stream from PyTorch+ cudaStream_t stream = at::cuda::getCurrentCUDAStream(a.device().index());++ nvfp4_gemv_batched_kernel<<<blocks, threads, smem_size, stream>>>(+ a.data_ptr<uint8_t>(),+ b.data_ptr<uint8_t>(),+ reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr<at::Float8_e4m3fn>()),+ reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr<at::Float8_e4m3fn>()),+ reinterpret_cast<half*>(c.data_ptr<at::Half>()),+ M, K, L+ );++ // Check for kernel launch errors+ AT_CUDA_CHECK(cudaGetLastError());+ }+ """++ cpp_source = """+ void nvfp4_gemv_cuda(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa,+ torch::Tensor sfb,+ torch::Tensor c,+ int M,+ int K+ );++ void nvfp4_gemv_batched_cuda(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa,+ torch::Tensor sfb,+ torch::Tensor c,+ int M,+ int K,+ int L+ );+ """++ # Compile the CUDA extension inline+ nvfp4_gemv_module = load_inline(+ name='nvfp4_gemv',+ cpp_sources=[cpp_source],+ cuda_sources=[cuda_source],+ functions=['nvfp4_gemv_cuda', 'nvfp4_gemv_batched_cuda'],+ verbose=True,+ extra_cuda_cflags=[+ '-O3',+ '--use_fast_math',+ '-arch=sm_100a',+ '--std=c++17',+ '-U__CUDA_NO_HALF_OPERATORS__', # Enable half operators+ '-U__CUDA_NO_HALF_CONVERSIONS__', # Enable half conversions+ ],+ )++ def custom_kernel(data: input_t) -> output_t:"""- PyTorch reference implementation of NVFP4 block-scaled GEMV.+ Custom CUDA implementation of NVFP4 block-scaled GEMV.+ Uses separate kernels optimized for L=1 and L>1 cases."""+ import torcha_ref, b_ref, sfa, sfb, _, _, c_ref = data- # a_ref is [m, k//2, l]- # b_ref is [n, k//2, l], n=1 padded to n=128- # c_ref is [m, 1, l]- return batched_gemv_impl(- a_ref,- b_ref,- c_ref,- sfa,- sfb,- )+ M, K_half, L = a_ref.shape+ K = K_half * 2 # Each byte contains 2 FP4 values+ if L == 1:+ # Use single-batch optimized kernel.+ # Pass c_ref directly - kernel writes c[m] which maps to c_ref[m, 0, 0].+ a_bytes = a_ref[:, :, 0].view(torch.uint8).contiguous()+ b_bytes = b_ref[0, :, 0].view(torch.uint8).contiguous()++ nvfp4_gemv_module.nvfp4_gemv_cuda(+ a_bytes,+ b_bytes,+ sfa[:, :, 0].contiguous(),+ sfb[0, :, 0].contiguous(),+ c_ref,+ M, K+ )+ else:+ # Use batched kernel for L>1 - processes all batches in one launch.+ # Original layout already has K contiguous (stride 1) from creation+ # as (L, M, K).permute(1, 2, 0), so no additional permute is needed.+ a_bytes = a_ref.view(torch.uint8)+ b_bytes = b_ref.view(torch.uint8)++ nvfp4_gemv_module.nvfp4_gemv_batched_cuda(+ a_bytes,+ b_bytes,+ sfa,+ sfb,+ c_ref,+ M, K, L+ )++ return c_ref
scrolls · 457 diff lines total
Best evidence level for this revision: reported
JSON