submission 77328
mdouglas · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 513 lines, June 9 Researcher Reciprocity License v1.0.
submission_cuda.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-77328?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:8756ff586ee9c43a1d2691ab3b4507f199c7e45cf5b6261fcb7bf37968bb801b
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
const uint4 a_data = *reinterpret_cast<const uint4*>(&a[a_base + k_byte_base]);Kernel source
submission_cuda.py513 lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import torch
# 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.
// Template parameter UseNestedLoops: true for large K (avoid div/mod), false for small K (less overhead)
template<bool UseNestedLoops>
__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;
const int N_padded = 128; // B is padded to 128 rows for torch._scaled_mm
// Cooperatively load all L B vectors into shared memory
// Template specialization: compile-time branch selection based on K size
// B original layout: b[n, k, l] at offset n*K_half + k + l*N_padded*K_half
// We only need n=0 (the actual vector, rest is padding)
if constexpr (UseNestedLoops) {
// Large K: nested loops to avoid expensive div/mod
for (int l = 0; l < L; l++) {
for (int k = tid; k < K_half; k += blockDim.x) {
sb[l * K_half + k] = b[k + l * N_padded * K_half];
}
}
} else {
// Small K: flat loop with div/mod has less overhead
for (int kl = tid; kl < K_half * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
sb[l * K_half + k] = b[k + l * N_padded * K_half];
}
}
// Cooperatively load scale factors for B
if constexpr (UseNestedLoops) {
// Large K: nested loops
for (int l = 0; l < L; l++) {
for (int k = tid; k < K_div_16; k += blockDim.x) {
ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
}
}
} else {
// Small K: flat loop
for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {
int k = kl / L;
int l = kl % L;
ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];
}
}
__syncthreads();
// Parallelize across L dimension: each warp handles one (M, L) pair
const int WARPS_PER_BLOCK = blockDim.x / 32;
const int WARPS_PER_M_ROW = L; // One warp per L batch
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;
int m_base = blockIdx.x * M_ROWS_PER_BLOCK;
// Which (M, L) pair does this warp handle?
int m_local = warp_id / WARPS_PER_M_ROW;
int l = warp_id % WARPS_PER_M_ROW;
int m = m_base + m_local;
if (m >= M) return;
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 2 scale blocks per iteration for better ILP (16 bytes with uint4)
int num_scale_pairs = K_div_16 / 2;
for (int scale_pair = lane; scale_pair < num_scale_pairs; scale_pair += 32) {
int scale_block_0 = scale_pair * 2;
int scale_block_1 = scale_block_0 + 1;
// Load both scale factors
half scale_a0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block_0].__x, __NV_E4M3).x);
half scale_b0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block_0].__x, __NV_E4M3).x);
half combined_scale0 = scale_a0 * scale_b0;
__half2 scale2_0 = __half2half2(combined_scale0);
half scale_a1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[sfa_base + scale_block_1].__x, __NV_E4M3).x);
half scale_b1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb_row[scale_block_1].__x, __NV_E4M3).x);
half combined_scale1 = scale_a1 * scale_b1;
__half2 scale2_1 = __half2half2(combined_scale1);
// Load 16 bytes at once using uint4
int k_byte_base = scale_block_0 * 8;
const uint4 a_data = *reinterpret_cast<const uint4*>(&a[a_base + k_byte_base]);
const uint4 b_data = *reinterpret_cast<const uint4*>(&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 first 8 bytes with scale_0
__half2 local_sum_0 = __float2half2_rn(0.0f);
#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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_0);
local_sum_0 = __hadd2(local_sum_0, scaled);
}
sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));
// Process second 8 bytes with scale_1
__half2 local_sum_1 = __float2half2_rn(0.0f);
#pragma unroll
for (int i = 8; i < 16; 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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_1);
local_sum_1 = __hadd2(local_sum_1, scaled);
}
sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
}
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
) {
// 4 warps per M row - more iterations per thread for better ILP and latency hiding
// 8 M rows per block to maximize work per block
const int WARPS_PER_M_ROW = 4;
const int WARPS_PER_BLOCK = blockDim.x / 32;
const int M_ROWS_PER_BLOCK = 8;
// Shared memory: B vector, scale factors B, and warp partial sums
extern __shared__ uint8_t smem[];
uint8_t* sb = smem;
__nv_fp8_e4m3* ssfb = reinterpret_cast<__nv_fp8_e4m3*>(sb + K/2);
float* warp_sums = reinterpret_cast<float*>(ssfb + K/16); // WARPS_PER_BLOCK partial sums
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 2 M rows
int m_base = blockIdx.x * M_ROWS_PER_BLOCK;
// Which M row and K chunk does this warp handle?
int m_local = warp_id / WARPS_PER_M_ROW;
int m = m_base + m_local;
if (m >= M) return;
int warp_in_m_group = warp_id % WARPS_PER_M_ROW;
// Process 2 scale blocks per iteration for better ILP
int scale_pairs_per_warp = (K_div_16 / WARPS_PER_M_ROW) / 2; // 1024 / 8 / 2 = 64 scale pairs per warp
int scale_pair_start = warp_in_m_group * scale_pairs_per_warp;
int scale_pair_end = scale_pair_start + scale_pairs_per_warp;
float sum = 0.0f;
// Loop over scale pairs - each iteration processes 16 bytes (2 scale blocks)
// With 8 warps: 64 scale pairs / 32 threads = 2 iterations per thread
for (int scale_pair = scale_pair_start + lane; scale_pair < scale_pair_end; scale_pair += 32) {
int scale_block_0 = scale_pair * 2;
int scale_block_1 = scale_block_0 + 1;
// Load both scale factors
half scale_a0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block_0].__x, __NV_E4M3).x);
half scale_b0 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block_0].__x, __NV_E4M3).x);
half combined_scale0 = scale_a0 * scale_b0;
__half2 scale2_0 = __half2half2(combined_scale0);
half scale_a1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(sfa[m * K_div_16 + scale_block_1].__x, __NV_E4M3).x);
half scale_b1 = __ushort_as_half(__nv_cvt_fp8_to_halfraw(ssfb[scale_block_1].__x, __NV_E4M3).x);
half combined_scale1 = scale_a1 * scale_b1;
__half2 scale2_1 = __half2half2(combined_scale1);
int k_byte_base = scale_block_0 * 8;
// Load 16 bytes at once using uint4
const uint4 a_data = *reinterpret_cast<const uint4*>(&a[m * K_half + k_byte_base]);
const uint4 b_data = *reinterpret_cast<const uint4*>(&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 first 8 bytes with scale_0 - accumulate in half2 first
__half2 local_sum_0 = __float2half2_rn(0.0f);
#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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_0);
local_sum_0 = __hadd2(local_sum_0, scaled);
}
sum += __half2float(__hadd(local_sum_0.x, local_sum_0.y));
// Process second 8 bytes with scale_1 - accumulate in half2 first
__half2 local_sum_1 = __float2half2_rn(0.0f);
#pragma unroll
for (int i = 8; i < 16; 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 scaled = __hmul2(__hmul2(a_vals, b_vals), scale2_1);
local_sum_1 = __hadd2(local_sum_1, scaled);
}
sum += __half2float(__hadd(local_sum_1.x, local_sum_1.y));
}
// Intra-warp reduction
sum = warp_reduce_sum(sum);
// Lane 0 of each warp writes its partial sum to shared memory
if (lane == 0) {
warp_sums[warp_id] = sum;
}
__syncthreads();
// Final reduction: each of the first M_ROWS_PER_BLOCK threads reduces one M row
if (tid < M_ROWS_PER_BLOCK) {
int m_write = m_base + tid;
if (m_write < M) {
// Reduce WARPS_PER_M_ROW partial sums for this M row
float final_sum = 0.0f;
int warp_start = tid * WARPS_PER_M_ROW;
#pragma unroll
for (int w = 0; w < WARPS_PER_M_ROW; w++) {
final_sum += warp_sums[warp_start + w];
}
c[m_write] = __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
) {
// 4 warps per M row, 8 M rows per block
const int WARPS_PER_M_ROW = 4;
const int M_ROWS_PER_BLOCK = 8;
const int WARPS_PER_BLOCK = WARPS_PER_M_ROW * M_ROWS_PER_BLOCK; // 32
const int threads = WARPS_PER_BLOCK * 32; // 1024 threads
const int blocks = (M + M_ROWS_PER_BLOCK - 1) / M_ROWS_PER_BLOCK;
// Shared memory: B vector + sfb + warp 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
) {
// Parallelize across L: each warp handles one (M, L) pair
// Use 32 warps to process more M rows per block, reducing redundant B+sfb loads
const int WARPS_PER_BLOCK = 32;
const int WARPS_PER_M_ROW = L;
const int M_ROWS_PER_BLOCK = WARPS_PER_BLOCK / WARPS_PER_M_ROW;
const int threads = WARPS_PER_BLOCK * 32; // 1024 threads
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());
// Choose template specialization based on K size to avoid runtime branching
if (K >= 4096) {
// Large K: use nested loops to avoid div/mod
nvfp4_gemv_batched_kernel<true><<<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
);
} else {
// Small K: use flat loop with div/mod
nvfp4_gemv_batched_kernel<false><<<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());
}
// Dispatch function - handles L=1 extraction and L=8 splitting in C++
void nvfp4_gemv_dispatch_cuda(
torch::Tensor a_ref, // [M, K//2, L] FP4
torch::Tensor b_ref, // [128, K//2, L] FP4
torch::Tensor sfa, // [M, K//16, L] FP8
torch::Tensor sfb, // [128, K//16, L] FP8
torch::Tensor c_ref // [M, 1, L] FP16
) {
int M = a_ref.size(0);
int K_half = a_ref.size(1);
int L = a_ref.size(2);
int K = K_half * 2;
if (L == 1) {
// L=1: Extract L=0 slice (all operations are views, no copying)
auto a_slice = a_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});
auto b_slice = b_ref.index({0, torch::indexing::Slice(), 0});
auto sfa_slice = sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});
auto sfb_slice = sfb.index({0, torch::indexing::Slice(), 0});
auto a_bytes = a_slice.view(torch::kUInt8).contiguous();
auto b_bytes = b_slice.view(torch::kUInt8).contiguous();
nvfp4_gemv_cuda(
a_bytes,
b_bytes,
sfa_slice.contiguous(),
sfb_slice.contiguous(),
c_ref,
M, K
);
} else if (L == 8) {
// L=8: Split into two L=4 calls
auto a_bytes = a_ref.view(torch::kUInt8);
auto b_bytes = b_ref.view(torch::kUInt8);
// PyTorch slicing preserves strides - no copying
nvfp4_gemv_batched_cuda(
a_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),
b_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),
sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),
sfb.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),
c_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),
M, K, 4
);
nvfp4_gemv_batched_cuda(
a_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),
b_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),
sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),
sfb.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),
c_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),
M, K, 4
);
} else {
// Other L: use batched kernel
auto a_bytes = a_ref.view(torch::kUInt8);
auto b_bytes = b_ref.view(torch::kUInt8);
nvfp4_gemv_batched_cuda(
a_bytes, b_bytes, sfa, sfb, c_ref,
M, K, L
);
}
}
"""
cpp_source = """
void nvfp4_gemv_dispatch_cuda(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c
);
"""
# Compile the CUDA extension inline
nvfp4_gemv_module = load_inline(
name='nvfp4_gemv',
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
functions=['nvfp4_gemv_dispatch_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.
Dispatch logic now in C++ to minimize Python overhead.
"""
a_ref, b_ref, sfa, sfb, _, _, c_ref = data
# Single C++ call - dispatch logic handled in C++
nvfp4_gemv_module.nvfp4_gemv_dispatch_cuda(
a_ref,
b_ref,
sfa,
sfb,
c_ref
)
return c_ref
scrolls · 513 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 77181.
⋯ 28 unchanged lines// 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.+ // Template parameter UseNestedLoops: true for large K (avoid div/mod), false for small K (less overhead)+ template<bool UseNestedLoops>__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⋯ 19 unchanged linesconst 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)+ // Template specialization: compile-time branch selection based on K size+ // B original layout: b[n, k, l] at offset n*K_half + k + l*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];+ if constexpr (UseNestedLoops) {+ // Large K: nested loops to avoid expensive div/mod+ for (int l = 0; l < L; l++) {+ for (int k = tid; k < K_half; k += blockDim.x) {+ sb[l * K_half + k] = b[k + l * N_padded * K_half];+ }+ }+ } else {+ // Small K: flat loop with div/mod has less overhead+ for (int kl = tid; kl < K_half * L; kl += blockDim.x) {+ int k = kl / L;+ int l = kl % L;+ 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];+ if constexpr (UseNestedLoops) {+ // Large K: nested loops+ for (int l = 0; l < L; l++) {+ for (int k = tid; k < K_div_16; k += blockDim.x) {+ ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];+ }+ }+ } else {+ // Small K: flat loop+ for (int kl = tid; kl < K_div_16 * L; kl += blockDim.x) {+ int k = kl / L;+ int l = kl % L;+ ssfb[l * K_div_16 + k] = sfb[k + l * N_padded * K_div_16];+ }}__syncthreads();⋯ 14 unchanged linesfloat sum = 0.0f;- // K dimension has stride 1, enabling coalesced access.+ // K dimension has stride 1, enabling coalesced accessconst 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];⋯ 255 unchanged lines// Get current CUDA stream from PyTorchcudaStream_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- );+ // Choose template specialization based on K size to avoid runtime branching+ if (K >= 4096) {+ // Large K: use nested loops to avoid div/mod+ nvfp4_gemv_batched_kernel<true><<<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+ );+ } else {+ // Small K: use flat loop with div/mod+ nvfp4_gemv_batched_kernel<false><<<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 errorsAT_CUDA_CHECK(cudaGetLastError());}++ // Dispatch function - handles L=1 extraction and L=8 splitting in C+++ void nvfp4_gemv_dispatch_cuda(+ torch::Tensor a_ref, // [M, K//2, L] FP4+ torch::Tensor b_ref, // [128, K//2, L] FP4+ torch::Tensor sfa, // [M, K//16, L] FP8+ torch::Tensor sfb, // [128, K//16, L] FP8+ torch::Tensor c_ref // [M, 1, L] FP16+ ) {+ int M = a_ref.size(0);+ int K_half = a_ref.size(1);+ int L = a_ref.size(2);+ int K = K_half * 2;++ if (L == 1) {+ // L=1: Extract L=0 slice (all operations are views, no copying)+ auto a_slice = a_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});+ auto b_slice = b_ref.index({0, torch::indexing::Slice(), 0});+ auto sfa_slice = sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), 0});+ auto sfb_slice = sfb.index({0, torch::indexing::Slice(), 0});++ auto a_bytes = a_slice.view(torch::kUInt8).contiguous();+ auto b_bytes = b_slice.view(torch::kUInt8).contiguous();++ nvfp4_gemv_cuda(+ a_bytes,+ b_bytes,+ sfa_slice.contiguous(),+ sfb_slice.contiguous(),+ c_ref,+ M, K+ );+ } else if (L == 8) {+ // L=8: Split into two L=4 calls+ auto a_bytes = a_ref.view(torch::kUInt8);+ auto b_bytes = b_ref.view(torch::kUInt8);++ // PyTorch slicing preserves strides - no copying+ nvfp4_gemv_batched_cuda(+ a_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),+ b_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),+ sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),+ sfb.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),+ c_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(0,4)}),+ M, K, 4+ );++ nvfp4_gemv_batched_cuda(+ a_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),+ b_bytes.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),+ sfa.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),+ sfb.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),+ c_ref.index({torch::indexing::Slice(), torch::indexing::Slice(), torch::indexing::Slice(4,8)}),+ M, K, 4+ );+ } else {+ // Other L: use batched kernel+ auto a_bytes = a_ref.view(torch::kUInt8);+ auto b_bytes = b_ref.view(torch::kUInt8);++ nvfp4_gemv_batched_cuda(+ a_bytes, b_bytes, sfa, sfb, c_ref,+ M, K, L+ );+ }+ }"""cpp_source = """- void nvfp4_gemv_cuda(+ void nvfp4_gemv_dispatch_cuda(torch::Tensor a,torch::Tensor b,torch::Tensor sfa,torch::Tensor sfb,- torch::Tensor c,- int M,- int K+ torch::Tensor c);-- 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⋯ 1 unchanged linesname='nvfp4_gemv',cpp_sources=[cpp_source],cuda_sources=[cuda_source],- functions=['nvfp4_gemv_cuda', 'nvfp4_gemv_batched_cuda'],+ functions=['nvfp4_gemv_dispatch_cuda'],verbose=True,extra_cuda_cflags=['-O3',⋯ 8 unchanged linesdef 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.+ Dispatch logic now in C++ to minimize Python overhead."""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+ # Single C++ call - dispatch logic handled in C+++ nvfp4_gemv_module.nvfp4_gemv_dispatch_cuda(+ a_ref,+ b_ref,+ sfa,+ sfb,+ c_ref+ )- 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 · 262 diff lines total
Best evidence level for this revision: reported
JSON