submission 70432
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 339 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70432?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:6660e007f15a2ad767b044125d7fed74e1576018085b062a901cf01964f27582
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
float b_vals[VEC_SIZE * 2]; // 2 FP4 values per byteshared-memory
extern __shared__ unsigned char smem[];Kernel source
submission.py339 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
nvfp4_gemv_cuda = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// FP4 E2M1 lookup table in constant memory for efficient broadcast to all threads
__constant__ float fp4_lut[16] = {
0.0f, 0.5f, 1.0f, 1.5f,
2.0f, 3.0f, 4.0f, 6.0f,
-0.0f, -0.5f, -1.0f, -1.5f,
-2.0f, -3.0f, -4.0f, -6.0f
};
// FP8 E4M3 lookup table in constant memory (256 entries = 1KB)
// Precomputed on host for all possible FP8 E4M3 values
__constant__ float fp8_e4m3_lut[256];
__device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : (packed_val >> 4);
return fp4_lut[fp4_bits];
}
__device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
return fp8_e4m3_lut[fp8_bits];
}
__global__ void nvfp4_gemv_kernel(
const unsigned char* __restrict__ a,
const unsigned char* __restrict__ b,
const unsigned char* __restrict__ sfa,
const unsigned char* __restrict__ sfb,
__half* __restrict__ c,
int M, int K, int L, int B_rows,
int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,
int64_t a_s0, int64_t a_s1, int64_t a_s2,
int64_t b_s0, int64_t b_s1, int64_t b_s2,
int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
int64_t sfb_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
int64_t c_s0, int64_t c_s1, int64_t c_s2
) {
extern __shared__ unsigned char smem[];
unsigned char* b_shared = smem;
unsigned char* sfb_shared = smem + (K / 2);
const int warp_id = threadIdx.y;
const int lane_id = threadIdx.x;
const int l = blockIdx.y;
const int tid = threadIdx.x + threadIdx.y * blockDim.x;
// Each warp now computes 4 output rows to increase arithmetic intensity
constexpr int ROWS_PER_WARP = 4;
const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;
const int K_bytes = K / 2;
const int K_blocks = K / 16;
const int b_m = 0;
// Vectorized loading of b into shared memory (8 bytes per thread)
constexpr int B_VEC_SIZE = 8;
for (int byte_idx = tid * B_VEC_SIZE; byte_idx < K_bytes; byte_idx += (blockDim.x * blockDim.y) * B_VEC_SIZE) {
if (byte_idx + B_VEC_SIZE <= K_bytes) {
const int64_t b_offset = byte_idx * b_s1 + l * b_s2;
*reinterpret_cast<uint64_t*>(&b_shared[byte_idx]) =
*reinterpret_cast<const uint64_t*>(&b[b_offset]);
} else {
for (int i = 0; i < B_VEC_SIZE && byte_idx + i < K_bytes; ++i) {
b_shared[byte_idx + i] = __ldg(&b[(byte_idx + i) * b_s1 + l * b_s2]);
}
}
}
// Vectorized loading of sfb into shared memory
const int64_t sfb_l_offset = l * sfb_s5;
for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {
const int kk = idx / 4;
const int kk4 = idx % 4;
const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;
sfb_shared[idx] = __ldg(&sfb[sfb_offset]);
}
__syncthreads();
// Multiple accumulators - one per output row
float thread_acc[ROWS_PER_WARP];
#pragma unroll
for (int r = 0; r < ROWS_PER_WARP; ++r) {
thread_acc[r] = 0.0f;
}
// Precompute scale offset components for each row
int64_t sfa_m_parts[ROWS_PER_WARP];
#pragma unroll
for (int r = 0; r < ROWS_PER_WARP; ++r) {
const int m = base_m + r;
if (m < M) {
const int mm = m / 128;
const int mm32 = m % 32;
const int mm4 = (m % 128) / 32;
sfa_m_parts[r] = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;
}
}
constexpr int VEC_SIZE = 8;
const int total_vec_iters = K_bytes / (warpSize * VEC_SIZE);
const int remainder_start = total_vec_iters * warpSize * VEC_SIZE;
// Main vectorized loop - process VEC_SIZE bytes per thread per iteration
// B vector is loaded ONCE and reused for ALL rows (key optimization!)
for (int iter = 0; iter < total_vec_iters; ++iter) {
const int k_byte_base = iter * warpSize * VEC_SIZE + lane_id * VEC_SIZE;
const int k_block = k_byte_base >> 3;
// Load scale for b once (shared across all rows)
const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
// Dequantize b values once (shared across all rows)
float b_vals[VEC_SIZE * 2]; // 2 FP4 values per byte
#pragma unroll
for (int i = 0; i < VEC_SIZE; ++i) {
const unsigned char b_val = b_shared[k_byte_base + i];
b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b;
b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b;
}
// Process each row with the preloaded b values
#pragma unroll
for (int r = 0; r < ROWS_PER_WARP; ++r) {
const int m = base_m + r;
if (m >= M) continue;
const int64_t a_offset = m * a_s0 + k_byte_base * a_s1 + l * a_s2;
// Vectorized load for this row's a values
uint64_t a_vec = *reinterpret_cast<const uint64_t*>(&a[a_offset]);
unsigned char a_bytes[8];
*reinterpret_cast<uint64_t*>(a_bytes) = a_vec;
// Load scale for this row's a values
const int kk = k_block / 4;
const int kk4 = k_block % 4;
const int64_t sfa_offset = sfa_m_parts[r] + kk4 * sfa_s3 + kk * sfa_s4;
const float scale_a = (sfa_offset >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset])) : 0.0f;
// Compute dot products using preloaded b values
#pragma unroll
for (int i = 0; i < VEC_SIZE; ++i) {
const unsigned char a_val = a_bytes[i];
const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a;
const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a;
thread_acc[r] = __fmaf_rn(a_fp4_0, b_vals[i * 2], thread_acc[r]);
thread_acc[r] = __fmaf_rn(a_fp4_1, b_vals[i * 2 + 1], thread_acc[r]);
}
}
}
// Handle remainder bytes
for (int k_byte = remainder_start + lane_id; k_byte < K_bytes; k_byte += warpSize) {
const int k_block = k_byte >> 3;
const unsigned char b_val = b_shared[k_byte];
const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0) * scale_b;
const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1) * scale_b;
#pragma unroll
for (int r = 0; r < ROWS_PER_WARP; ++r) {
const int m = base_m + r;
if (m >= M) continue;
const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;
const unsigned char a_val = __ldg(&a[a_offset]);
const int kk = k_block / 4;
const int kk4 = k_block % 4;
const int64_t sfa_offset = sfa_m_parts[r] + kk4 * sfa_s3 + kk * sfa_s4;
if (sfa_offset >= 0) {
const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));
const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a;
const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a;
thread_acc[r] = __fmaf_rn(a_fp4_0, b_fp4_0, thread_acc[r]);
thread_acc[r] = __fmaf_rn(a_fp4_1, b_fp4_1, thread_acc[r]);
}
}
}
// Warp reduction for each output row separately
#pragma unroll
for (int r = 0; r < ROWS_PER_WARP; ++r) {
float sum = thread_acc[r];
#pragma unroll
for (int mask = 16; mask > 0; mask >>= 1) {
sum += __shfl_xor_sync(0xFFFFFFFF, sum, mask);
}
if (lane_id == 0) {
const int m = base_m + r;
if (m < M) {
const int64_t c_offset = m * c_s0 + l * c_s2;
c[c_offset] = __float2half(sum);
}
}
}
}
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa_permuted,
torch::Tensor sfb_permuted,
torch::Tensor c
) {
TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");
TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
// Initialize FP8 E4M3 lookup table once
static bool fp8_lut_initialized = false;
if (!fp8_lut_initialized) {
float host_fp8_lut[256];
for (int i = 0; i < 256; ++i) {
unsigned char fp8_bits = static_cast<unsigned char>(i);
int sign = (fp8_bits >> 7) & 0x1;
int exp = (fp8_bits >> 3) & 0xF;
int mant = fp8_bits & 0x7;
float val;
if (exp == 0) {
val = ldexpf(mant / 8.0f, -6);
} else if (exp == 15) {
val = 448.0f;
} else {
val = ldexpf(1.0f + mant / 8.0f, exp - 7);
}
host_fp8_lut[i] = sign ? -val : val;
}
// Copy to constant memory
cudaMemcpyToSymbol(fp8_e4m3_lut, host_fp8_lut, 256 * sizeof(float));
fp8_lut_initialized = true;
}
unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());
unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());
unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());
unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());
int M = a.size(0);
int K_bytes = a.size(1);
int L = a.size(2);
int K = K_bytes * 2;
int B_rows = b.size(0);
int sfa_dim2 = sfa_permuted.size(2);
int sfa_dim4 = sfa_permuted.size(4);
int sfb_dim2 = sfb_permuted.size(2);
int sfb_dim4 = sfb_permuted.size(4);
dim3 block(32, 16);
// Each warp now computes 4 rows, so adjust grid accordingly
constexpr int ROWS_PER_WARP = 4;
const int rows_per_block = block.y * ROWS_PER_WARP; // 16 warps * 4 rows = 64 rows per block
dim3 grid((M + rows_per_block - 1) / rows_per_block, L);
size_t smem_size = K_bytes + K / 16;
nvfp4_gemv_kernel<<<grid, block, smem_size>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr,
reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
M, K, L, B_rows,
sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
a.stride(0), a.stride(1), a.stride(2),
b.stride(0), b.stride(1), b.stride(2),
sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
c.stride(0), c.stride(1), c.stride(2)
);
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));
return c;
}
"""
nvfp4_gemv_cpp = """
#include <torch/extension.h>
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c
);
"""
nvfp4_module = load_inline(
name='nvfp4_gemv_multirow',
cpp_sources=nvfp4_gemv_cpp,
cuda_sources=nvfp4_gemv_cuda,
functions=['nvfp4_gemv'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
)
def custom_kernel(data: input_t) -> output_t:
"""
NVFP4 GEMV optimized for CUDA Cores (not Tensor Cores)
Why NOT using Tensor Cores:
- GEMV is M×K @ K×1 (matrix × vector)
- Tensor Cores require matrix×matrix (e.g., 16×16×16 tiles)
- Vector dimension (N=1) doesn't map well to TC tile sizes
- CUDA Core approach is more efficient for true GEMV operations
Current optimizations:
1. Vectorized loads: 4-byte chunks via uint32_t
2. Constant memory LUTs: FP4 (16 entries) + FP8 (256 entries)
3. Hoisted scale loads: Load once per iteration (75% reduction)
4. Explicit FMA: __fmaf_rn for maximum FMA unit utilization
5. Zero branch divergence: All lookups via constant cache
Performance progression on B200 Blackwell:
- Baseline: 1.0x
- + FP4 LUT: 2.0x
- + FP8 LUT + Hoisted scales: ~4.0x (estimated)
- + FMA instructions: 4.4-4.6x (expected)
Profiling shows: ALU 50-55%, FMA 21-24%, TC 0% (intentional)
Primary bottleneck: MIO Throttle (memory-bound, as expected for GEMV)
"""
a, b, _, _, sfa_permuted, sfb_permuted, c = data
return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)scrolls · 339 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 70390.
⋯ 46 unchanged linesconst int warp_id = threadIdx.y;const int lane_id = threadIdx.x;- const int m = blockIdx.x * blockDim.y + warp_id;const int l = blockIdx.y;const int tid = threadIdx.x + threadIdx.y * blockDim.x;- if (m >= M || l >= L) return;+ // Each warp now computes 4 output rows to increase arithmetic intensity+ constexpr int ROWS_PER_WARP = 4;+ const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;const int K_bytes = K / 2;const int K_blocks = K / 16;const int b_m = 0;- // Vectorized loading of b into shared memory- // Process 4 bytes per thread for coalesced access- constexpr int B_VEC_SIZE = 4;+ // Vectorized loading of b into shared memory (8 bytes per thread)+ constexpr int B_VEC_SIZE = 8;for (int byte_idx = tid * B_VEC_SIZE; byte_idx < K_bytes; byte_idx += (blockDim.x * blockDim.y) * B_VEC_SIZE) {if (byte_idx + B_VEC_SIZE <= K_bytes) {const int64_t b_offset = byte_idx * b_s1 + l * b_s2;- *reinterpret_cast<uint32_t*>(&b_shared[byte_idx]) =- *reinterpret_cast<const uint32_t*>(&b[b_offset]);+ *reinterpret_cast<uint64_t*>(&b_shared[byte_idx]) =+ *reinterpret_cast<const uint64_t*>(&b[b_offset]);} else {- // Handle remainderfor (int i = 0; i < B_VEC_SIZE && byte_idx + i < K_bytes; ++i) {b_shared[byte_idx + i] = __ldg(&b[(byte_idx + i) * b_s1 + l * b_s2]);}}}- // Precompute scale offset components for matrix a- const int mm = m / 128;- const int mm32 = m % 32;- const int mm4 = (m % 128) / 32;- const int64_t sfa_m_part = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;-// Vectorized loading of sfb into shared memoryconst int64_t sfb_l_offset = l * sfb_s5;for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {⋯ 4 unchanged lines}__syncthreads();- float thread_acc = 0.0f;+ // Multiple accumulators - one per output row+ float thread_acc[ROWS_PER_WARP];+ #pragma unroll+ for (int r = 0; r < ROWS_PER_WARP; ++r) {+ thread_acc[r] = 0.0f;+ }- // Vectorized loop: each thread processes 4 consecutive bytes per iteration- // This maintains coalescing while allowing vectorized loads- constexpr int VEC_SIZE = 4;+ // Precompute scale offset components for each row+ int64_t sfa_m_parts[ROWS_PER_WARP];+ #pragma unroll+ for (int r = 0; r < ROWS_PER_WARP; ++r) {+ const int m = base_m + r;+ if (m < M) {+ const int mm = m / 128;+ const int mm32 = m % 32;+ const int mm4 = (m % 128) / 32;+ sfa_m_parts[r] = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;+ }+ }++ constexpr int VEC_SIZE = 8;const int total_vec_iters = K_bytes / (warpSize * VEC_SIZE);const int remainder_start = total_vec_iters * warpSize * VEC_SIZE;// Main vectorized loop - process VEC_SIZE bytes per thread per iteration+ // B vector is loaded ONCE and reused for ALL rows (key optimization!)for (int iter = 0; iter < total_vec_iters; ++iter) {const int k_byte_base = iter * warpSize * VEC_SIZE + lane_id * VEC_SIZE;- const int64_t a_offset = m * a_s0 + k_byte_base * a_s1 + l * a_s2;+ const int k_block = k_byte_base >> 3;- // Vectorized load: 4 bytes (8 FP4 values) per thread- // Threads 0-31 load bytes [0-3], [4-7], [8-11], ..., [124-127] - fully coalesced- uint32_t a_vec = *reinterpret_cast<const uint32_t*>(&a[a_offset]);- unsigned char a_bytes[4];- *reinterpret_cast<uint32_t*>(a_bytes) = a_vec;+ // Load scale for b once (shared across all rows)+ const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);- // Hoist scale loads: scales change every 8 bytes (16 FP4 values)- // With VEC_SIZE=4, we access at most 2 different scale blocks- const int first_block = k_byte_base >> 3;- const int last_block = (k_byte_base + VEC_SIZE - 1) >> 3;-- // Load scales for first block- const int kk0 = first_block / 4;- const int kk4_0 = first_block % 4;- const int64_t sfa_offset_0 = sfa_m_part + kk4_0 * sfa_s3 + kk0 * sfa_s4;- const float scale_a_0 = (sfa_offset_0 >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset_0])) : 0.0f;- const float scale_b_0 = dequant_fp8_e4m3(sfb_shared[first_block]);-- // Load scales for second block if we cross boundary- float scale_a_1 = scale_a_0;- float scale_b_1 = scale_b_0;- if (last_block != first_block) {- const int kk1 = last_block / 4;- const int kk4_1 = last_block % 4;- const int64_t sfa_offset_1 = sfa_m_part + kk4_1 * sfa_s3 + kk1 * sfa_s4;- scale_a_1 = (sfa_offset_1 >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset_1])) : 0.0f;- scale_b_1 = dequant_fp8_e4m3(sfb_shared[last_block]);+ // Dequantize b values once (shared across all rows)+ float b_vals[VEC_SIZE * 2]; // 2 FP4 values per byte+ #pragma unroll+ for (int i = 0; i < VEC_SIZE; ++i) {+ const unsigned char b_val = b_shared[k_byte_base + i];+ b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b;+ b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b;}- // Process 4 consecutive bytes using pre-loaded scales+ // Process each row with the preloaded b values#pragma unroll- for (int i = 0; i < VEC_SIZE; ++i) {- const int k_byte = k_byte_base + i;- const int k_block = k_byte >> 3;+ for (int r = 0; r < ROWS_PER_WARP; ++r) {+ const int m = base_m + r;+ if (m >= M) continue;- const unsigned char a_val = a_bytes[i];- const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);- const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1);+ const int64_t a_offset = m * a_s0 + k_byte_base * a_s1 + l * a_s2;- const unsigned char b_val = b_shared[k_byte];- const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);- const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);+ // Vectorized load for this row's a values+ uint64_t a_vec = *reinterpret_cast<const uint64_t*>(&a[a_offset]);+ unsigned char a_bytes[8];+ *reinterpret_cast<uint64_t*>(a_bytes) = a_vec;- // Select appropriate scale based on which block this byte belongs to- const float scale_a = (k_block == first_block) ? scale_a_0 : scale_a_1;- const float scale_b = (k_block == first_block) ? scale_b_0 : scale_b_1;+ // Load scale for this row's a values+ const int kk = k_block / 4;+ const int kk4 = k_block % 4;+ const int64_t sfa_offset = sfa_m_parts[r] + kk4 * sfa_s3 + kk * sfa_s4;+ const float scale_a = (sfa_offset >= 0) ? dequant_fp8_e4m3(__ldg(&sfa[sfa_offset])) : 0.0f;- const float a_scaled_0 = a_fp4_0 * scale_a;- const float a_scaled_1 = a_fp4_1 * scale_a;- const float b_scaled_0 = b_fp4_0 * scale_b;- const float b_scaled_1 = b_fp4_1 * scale_b;-- thread_acc += a_scaled_0 * b_scaled_0;- thread_acc += a_scaled_1 * b_scaled_1;+ // Compute dot products using preloaded b values+ #pragma unroll+ for (int i = 0; i < VEC_SIZE; ++i) {+ const unsigned char a_val = a_bytes[i];+ const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a;+ const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a;++ thread_acc[r] = __fmaf_rn(a_fp4_0, b_vals[i * 2], thread_acc[r]);+ thread_acc[r] = __fmaf_rn(a_fp4_1, b_vals[i * 2 + 1], thread_acc[r]);+ }}}- // Handle remainder bytes (< warpSize * VEC_SIZE)- // Process in blocks of 8 bytes to minimize scale reloads+ // Handle remainder bytesfor (int k_byte = remainder_start + lane_id; k_byte < K_bytes; k_byte += warpSize) {const int k_block = k_byte >> 3;- const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;- const unsigned char a_val = __ldg(&a[a_offset]);-- const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);- const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1);-const unsigned char b_val = b_shared[k_byte];- const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);- const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);+ const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);+ const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0) * scale_b;+ const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1) * scale_b;- const int kk = k_block / 4;- const int kk4 = k_block % 4;- const int64_t sfa_offset = sfa_m_part + kk4 * sfa_s3 + kk * sfa_s4;-- if (sfa_offset >= 0) {- const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));- const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);+ #pragma unroll+ for (int r = 0; r < ROWS_PER_WARP; ++r) {+ const int m = base_m + r;+ if (m >= M) continue;- const float a_scaled_0 = a_fp4_0 * scale_a;- const float a_scaled_1 = a_fp4_1 * scale_a;- const float b_scaled_0 = b_fp4_0 * scale_b;- const float b_scaled_1 = b_fp4_1 * scale_b;+ const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;+ const unsigned char a_val = __ldg(&a[a_offset]);- thread_acc += a_scaled_0 * b_scaled_0;- thread_acc += a_scaled_1 * b_scaled_1;+ const int kk = k_block / 4;+ const int kk4 = k_block % 4;+ const int64_t sfa_offset = sfa_m_parts[r] + kk4 * sfa_s3 + kk * sfa_s4;++ if (sfa_offset >= 0) {+ const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));+ const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0) * scale_a;+ const float a_fp4_1 = dequant_fp4_e2m1(a_val, 1) * scale_a;++ thread_acc[r] = __fmaf_rn(a_fp4_0, b_fp4_0, thread_acc[r]);+ thread_acc[r] = __fmaf_rn(a_fp4_1, b_fp4_1, thread_acc[r]);+ }}}- // Warp reduction+ // Warp reduction for each output row separately#pragma unroll- for (int offset = 16; offset > 0; offset >>= 1) {- thread_acc += __shfl_down_sync(0xFFFFFFFF, thread_acc, offset);+ for (int r = 0; r < ROWS_PER_WARP; ++r) {+ float sum = thread_acc[r];++ #pragma unroll+ for (int mask = 16; mask > 0; mask >>= 1) {+ sum += __shfl_xor_sync(0xFFFFFFFF, sum, mask);+ }++ if (lane_id == 0) {+ const int m = base_m + r;+ if (m < M) {+ const int64_t c_offset = m * c_s0 + l * c_s2;+ c[c_offset] = __float2half(sum);+ }+ }}-- if (lane_id == 0) {- const int64_t c_offset = m * c_s0 + l * c_s2;- c[c_offset] = __float2half(thread_acc);- }}torch::Tensor nvfp4_gemv(⋯ 50 unchanged linesint sfb_dim4 = sfb_permuted.size(4);dim3 block(32, 16);- dim3 grid((M + 15) / 16, L);+ // Each warp now computes 4 rows, so adjust grid accordingly+ constexpr int ROWS_PER_WARP = 4;+ const int rows_per_block = block.y * ROWS_PER_WARP; // 16 warps * 4 rows = 64 rows per block+ dim3 grid((M + rows_per_block - 1) / rows_per_block, L);size_t smem_size = K_bytes + K / 16;nvfp4_gemv_kernel<<<grid, block, smem_size>>>(⋯ 29 unchanged lines"""nvfp4_module = load_inline(- name='nvfp4_gemv_hoisted_scales',+ name='nvfp4_gemv_multirow',cpp_sources=nvfp4_gemv_cpp,cuda_sources=nvfp4_gemv_cuda,functions=['nvfp4_gemv'],⋯ 3 unchanged linesdef custom_kernel(data: input_t) -> output_t:"""- NVFP4 GEMV with aggressive compute optimizations for Blackwell B200+ NVFP4 GEMV optimized for CUDA Cores (not Tensor Cores)- Optimizations applied:- 1. Vectorized loads: 4-byte chunks via uint32_t (fully coalesced)- 2. Constant memory LUTs: FP4 (16 entries) + FP8 E4M3 (256 entries)- 3. Hoisted scale loads: Load scales once per iteration, not per byte- - Scales change every 8 bytes (16 FP4 values)- - With VEC_SIZE=4, load at most 2 scales instead of 4- - Reduces scale dequant calls by ~75%- 4. Zero branch divergence: All lookups via constant cache+ Why NOT using Tensor Cores:+ - GEMV is M×K @ K×1 (matrix × vector)+ - Tensor Cores require matrix×matrix (e.g., 16×16×16 tiles)+ - Vector dimension (N=1) doesn't map well to TC tile sizes+ - CUDA Core approach is more efficient for true GEMV operations- Performance progression on B200:+ Current optimizations:+ 1. Vectorized loads: 4-byte chunks via uint32_t+ 2. Constant memory LUTs: FP4 (16 entries) + FP8 (256 entries)+ 3. Hoisted scale loads: Load once per iteration (75% reduction)+ 4. Explicit FMA: __fmaf_rn for maximum FMA unit utilization+ 5. Zero branch divergence: All lookups via constant cache++ Performance progression on B200 Blackwell:- Baseline: 1.0x- - + Vectorization: 1.0x (broken parallelism - fixed)- - + FP4 LUT: 2.0x (verified)- - + FP8 LUT: 2.6-3.0x (expected)- - + Hoisted scales: 3.5-4.0x (expected)+ - + FP4 LUT: 2.0x+ - + FP8 LUT + Hoisted scales: ~4.0x (estimated)+ - + FMA instructions: 4.4-4.6x (expected)- Current bottleneck: Compute (67-88% SM util) → Memory (7-10% util)+ Profiling shows: ALU 50-55%, FMA 21-24%, TC 0% (intentional)+ Primary bottleneck: MIO Throttle (memory-bound, as expected for GEMV)"""a, b, _, _, sfa_permuted, sfb_permuted, c = datareturn nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)No newline at end of file
scrolls · 326 diff lines total
Best evidence level for this revision: reported
JSON