submission 107169
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 199 lines, June 9 Researcher Reciprocity License v1.0.
kmajor.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-107169?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:54bf3fa35db9b9386de3f81793126c7c95083541653c68b7da3eb65413658714
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
__nv_fp8_e4m3 v; v.__x = x;vector-width = half2
half2* a_h2_0 = reinterpret_cast<half2*>(&a_f16_0);Kernel source
kmajor.py199 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_source = """
#include <cuda_fp16.h>
#include <cuda_fp8.h>
__device__ __forceinline__ float fp8_to_float(uint8_t x) {
__nv_fp8_e4m3 v; v.__x = x;
return float(v);
}
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xffffffff, val, offset);
return val;
}
// Convert 8 FP4 values (packed in 1 uint32) to 8 FP16 values (in 2 uint32s)
// Uses PTX cvt.rn.f16x2.e2m1x2 instruction
__device__ __forceinline__ void cvt_f4x8_to_f16x8(uint32_t src, uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3) {
asm volatile(
"{ .reg .b8 b0, b1, b2, b3; "
"mov.b32 {b0, b1, b2, b3}, %4; "
"cvt.rn.f16x2.e2m1x2 %0, b0; "
"cvt.rn.f16x2.e2m1x2 %1, b1; "
"cvt.rn.f16x2.e2m1x2 %2, b2; "
"cvt.rn.f16x2.e2m1x2 %3, b3; }"
: "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)
: "r"(src)
);
}
// V5: Native PTX FP4->FP16 conversion
__global__ void gemv_v5(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
half* __restrict__ C,
int M, int K, int L,
int64_t a_s0, int64_t a_s2,
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_s3, int64_t sfb_s4, int64_t sfb_s5,
int64_t c_s0, int64_t c_s2
) {
const int WARPS_PER_BLOCK = 8;
int warp_id = threadIdx.x / 32;
int lane_id = threadIdx.x % 32;
int row = blockIdx.x * WARPS_PER_BLOCK + warp_id;
int batch = blockIdx.y;
if (row >= M) return;
// Precompute base addresses
int64_t sfa_row_base = (row & 31) * sfa_s0 + ((row & 127) >> 5) * sfa_s1 +
(row / 128) * sfa_s2 + batch * sfa_s5;
int64_t sfb_batch_base = batch * sfb_s5;
const uint8_t* A_row = A + row * a_s0 + batch * a_s2;
const uint8_t* B_batch = B + batch * b_s2;
float acc = 0.0f;
int K_scales = K / 16;
// Process 2 scale groups (32 elements) per iteration using 128-bit loads
for (int scale_base = 0; scale_base < K_scales; scale_base += 32) {
int k_scale = scale_base + lane_id;
if (k_scale >= K_scales) break;
// Load scale factors
int sfa_idx = sfa_row_base + (k_scale & 3) * sfa_s3 + (k_scale >> 2) * sfa_s4;
int sfb_idx = sfb_batch_base + (k_scale & 3) * sfb_s3 + (k_scale >> 2) * sfb_s4;
float scale = fp8_to_float(SFA[sfa_idx]) * fp8_to_float(SFB[sfb_idx]);
// Load 8 bytes = 16 FP4 elements, but we only need 8 for one scale group
// Actually, let's load 4 bytes = 8 FP4 elements = half a scale group
// Wait - 16 elements per scale group, 8 bytes per scale group
// Let's load 8 bytes with uint2, process 16 elements
int k_byte_start = k_scale * 8;
uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte_start);
uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte_start);
// Convert first 4 bytes (8 FP4 values) of A
uint32_t a_f16_0, a_f16_1, a_f16_2, a_f16_3;
cvt_f4x8_to_f16x8(a_vec.x, a_f16_0, a_f16_1, a_f16_2, a_f16_3);
// Convert second 4 bytes (8 FP4 values) of A
uint32_t a_f16_4, a_f16_5, a_f16_6, a_f16_7;
cvt_f4x8_to_f16x8(a_vec.y, a_f16_4, a_f16_5, a_f16_6, a_f16_7);
// Convert first 4 bytes (8 FP4 values) of B
uint32_t b_f16_0, b_f16_1, b_f16_2, b_f16_3;
cvt_f4x8_to_f16x8(b_vec.x, b_f16_0, b_f16_1, b_f16_2, b_f16_3);
// Convert second 4 bytes (8 FP4 values) of B
uint32_t b_f16_4, b_f16_5, b_f16_6, b_f16_7;
cvt_f4x8_to_f16x8(b_vec.y, b_f16_4, b_f16_5, b_f16_6, b_f16_7);
// Compute dot product using half2 operations
half2* a_h2_0 = reinterpret_cast<half2*>(&a_f16_0);
half2* a_h2_1 = reinterpret_cast<half2*>(&a_f16_1);
half2* a_h2_2 = reinterpret_cast<half2*>(&a_f16_2);
half2* a_h2_3 = reinterpret_cast<half2*>(&a_f16_3);
half2* a_h2_4 = reinterpret_cast<half2*>(&a_f16_4);
half2* a_h2_5 = reinterpret_cast<half2*>(&a_f16_5);
half2* a_h2_6 = reinterpret_cast<half2*>(&a_f16_6);
half2* a_h2_7 = reinterpret_cast<half2*>(&a_f16_7);
half2* b_h2_0 = reinterpret_cast<half2*>(&b_f16_0);
half2* b_h2_1 = reinterpret_cast<half2*>(&b_f16_1);
half2* b_h2_2 = reinterpret_cast<half2*>(&b_f16_2);
half2* b_h2_3 = reinterpret_cast<half2*>(&b_f16_3);
half2* b_h2_4 = reinterpret_cast<half2*>(&b_f16_4);
half2* b_h2_5 = reinterpret_cast<half2*>(&b_f16_5);
half2* b_h2_6 = reinterpret_cast<half2*>(&b_f16_6);
half2* b_h2_7 = reinterpret_cast<half2*>(&b_f16_7);
// Multiply and accumulate
half2 prod0 = __hmul2(*a_h2_0, *b_h2_0);
half2 prod1 = __hmul2(*a_h2_1, *b_h2_1);
half2 prod2 = __hmul2(*a_h2_2, *b_h2_2);
half2 prod3 = __hmul2(*a_h2_3, *b_h2_3);
half2 prod4 = __hmul2(*a_h2_4, *b_h2_4);
half2 prod5 = __hmul2(*a_h2_5, *b_h2_5);
half2 prod6 = __hmul2(*a_h2_6, *b_h2_6);
half2 prod7 = __hmul2(*a_h2_7, *b_h2_7);
// Sum all products
half2 sum01 = __hadd2(prod0, prod1);
half2 sum23 = __hadd2(prod2, prod3);
half2 sum45 = __hadd2(prod4, prod5);
half2 sum67 = __hadd2(prod6, prod7);
half2 sum0123 = __hadd2(sum01, sum23);
half2 sum4567 = __hadd2(sum45, sum67);
half2 sum_all = __hadd2(sum0123, sum4567);
float local_sum = __half2float(sum_all.x) + __half2float(sum_all.y);
acc += local_sum * scale;
}
acc = warp_reduce_sum(acc);
if (lane_id == 0) {
C[row * c_s0 + batch * c_s2] = __float2half(acc);
}
}
torch::Tensor gemv_cuda(
torch::Tensor a, torch::Tensor b,
torch::Tensor sfa, torch::Tensor sfb,
torch::Tensor c
) {
int M = a.size(0), K = a.size(1) * 2, L = a.size(2);
const int WARPS_PER_BLOCK = 8;
dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
dim3 block(32 * WARPS_PER_BLOCK);
gemv_v5<<<grid, block>>>(
reinterpret_cast<const uint8_t*>(a.data_ptr()),
reinterpret_cast<const uint8_t*>(b.data_ptr()),
reinterpret_cast<const uint8_t*>(sfa.data_ptr()),
reinterpret_cast<const uint8_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(c.data_ptr()),
M, K, L,
a.stride(0), a.stride(2),
b.stride(2),
sfa.stride(0), sfa.stride(1), sfa.stride(2), sfa.stride(3), sfa.stride(4), sfa.stride(5),
sfb.stride(3), sfb.stride(4), sfb.stride(5),
c.stride(0), c.stride(2)
);
return c;
}
"""
cpp_source = """
torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""
module = load_inline(
name='gemv_v5',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['gemv_cuda'],
extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17', '--generate-code=arch=compute_100a,code=sm_100a'],
verbose=True
)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, sfa_perm, sfb_perm, c = data
return module.gemv_cuda(a, b, sfa_perm, sfb_perm, c)scrolls · 199 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 70699.
import torchfrom torch.utils.cpp_extension import load_inline+ from task import input_t, output_t- nvfp4_gemv_cuda = """+ cuda_source = """#include <cuda_fp16.h>- #include <cuda_runtime.h>+ #include <cuda_fp8.h>- __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- };+ __device__ __forceinline__ float fp8_to_float(uint8_t x) {+ __nv_fp8_e4m3 v; v.__x = x;+ return float(v);+ }- __constant__ float fp8_e4m3_lut[256];-- __device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {- return fp4_lut[(idx == 0) ? (packed_val & 0x0F) : (packed_val >> 4)];+ __device__ __forceinline__ float warp_reduce_sum(float val) {+ #pragma unroll+ for (int offset = 16; offset > 0; offset >>= 1)+ val += __shfl_xor_sync(0xffffffff, val, offset);+ return val;}- __device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {- return fp8_e4m3_lut[fp8_bits];+ // Convert 8 FP4 values (packed in 1 uint32) to 8 FP16 values (in 2 uint32s)+ // Uses PTX cvt.rn.f16x2.e2m1x2 instruction+ __device__ __forceinline__ void cvt_f4x8_to_f16x8(uint32_t src, uint32_t& dst0, uint32_t& dst1, uint32_t& dst2, uint32_t& dst3) {+ asm volatile(+ "{ .reg .b8 b0, b1, b2, b3; "+ "mov.b32 {b0, b1, b2, b3}, %4; "+ "cvt.rn.f16x2.e2m1x2 %0, b0; "+ "cvt.rn.f16x2.e2m1x2 %1, b1; "+ "cvt.rn.f16x2.e2m1x2 %2, b2; "+ "cvt.rn.f16x2.e2m1x2 %3, b3; }"+ : "=r"(dst0), "=r"(dst1), "=r"(dst2), "=r"(dst3)+ : "r"(src)+ );}- template<int ROWS_PER_WARP, int VEC_SIZE>- __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,+ // V5: Native PTX FP4->FP16 conversion+ __global__ void gemv_v5(+ const uint8_t* __restrict__ A,+ const uint8_t* __restrict__ B,+ const uint8_t* __restrict__ SFA,+ const uint8_t* __restrict__ SFB,+ half* __restrict__ C,+ int M, int K, int L,+ int64_t a_s0, int64_t a_s2,+ 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+ int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,+ int64_t c_s0, int64_t c_s2) {- extern __shared__ unsigned char smem[];- unsigned char* b_shared = smem;- unsigned char* sfb_shared = smem + (K / 2) + 4;+ const int WARPS_PER_BLOCK = 8;- 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;+ int warp_id = threadIdx.x / 32;+ int lane_id = threadIdx.x % 32;+ int row = blockIdx.x * WARPS_PER_BLOCK + warp_id;+ int batch = blockIdx.y;- const int base_m = (blockIdx.x * blockDim.y + warp_id) * ROWS_PER_WARP;+ if (row >= M) return;- const int K_bytes = K / 2;- const int K_blocks = K / 16;- const int b_m = 0;+ // Precompute base addresses+ int64_t sfa_row_base = (row & 31) * sfa_s0 + ((row & 127) >> 5) * sfa_s1 ++ (row / 128) * sfa_s2 + batch * sfa_s5;+ int64_t sfb_batch_base = batch * sfb_s5;- constexpr int B_VEC_SIZE = 16;- 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 && b_s1 == 1) {- const int64_t b_offset_base = byte_idx * b_s1 + l * b_s2;- *reinterpret_cast<uint4*>(&b_shared[byte_idx]) =- __ldg(reinterpret_cast<const uint4*>(&b[b_offset_base]));- } 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]);- }- }- }+ const uint8_t* A_row = A + row * a_s0 + batch * a_s2;+ const uint8_t* B_batch = B + batch * b_s2;- 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();+ float acc = 0.0f;+ int K_scales = K / 16;- float thread_acc[ROWS_PER_WARP];- #pragma unroll- for (int r = 0; r < ROWS_PER_WARP; ++r) thread_acc[r] = 0.0f;-- int64_t sfa_m_parts[ROWS_PER_WARP];- bool row_valid[ROWS_PER_WARP];- #pragma unroll- for (int r = 0; r < ROWS_PER_WARP; ++r) {- const int m = base_m + r;- row_valid[r] = (m < M && l < L);- if (row_valid[r]) {- 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;- }- }-- const int total_vec_iters = K_bytes / (warpSize * VEC_SIZE);- const int remainder_start = total_vec_iters * warpSize * VEC_SIZE;-- for (int iter = 0; iter < total_vec_iters; ++iter) {- const int k_byte_base = iter * warpSize * VEC_SIZE + lane_id * VEC_SIZE;+ // Process 2 scale groups (32 elements) per iteration using 128-bit loads+ for (int scale_base = 0; scale_base < K_scales; scale_base += 32) {+ int k_scale = scale_base + lane_id;+ if (k_scale >= K_scales) break;- // Process in chunks to reduce register pressure- constexpr int CHUNK_SIZE = VEC_SIZE / 2;+ // Load scale factors+ int sfa_idx = sfa_row_base + (k_scale & 3) * sfa_s3 + (k_scale >> 2) * sfa_s4;+ int sfb_idx = sfb_batch_base + (k_scale & 3) * sfb_s3 + (k_scale >> 2) * sfb_s4;+ float scale = fp8_to_float(SFA[sfa_idx]) * fp8_to_float(SFB[sfb_idx]);- for (int chunk = 0; chunk < 2; ++chunk) {- const int chunk_offset = chunk * CHUNK_SIZE;- const int block_0 = (k_byte_base + chunk_offset) >> 3;- const int block_1 = (k_byte_base + chunk_offset + 8) >> 3;-- const float scale_b_0 = dequant_fp8_e4m3(sfb_shared[block_0]);- const float scale_b_1 = dequant_fp8_e4m3(sfb_shared[block_1]);-- float b_vals[CHUNK_SIZE * 2];-- #pragma unroll- for (int i = 0; i < 8; ++i) {- const unsigned char b_val = b_shared[k_byte_base + chunk_offset + i];- b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b_0;- b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b_0;- }- #pragma unroll- for (int i = 8; i < CHUNK_SIZE; ++i) {- const unsigned char b_val = b_shared[k_byte_base + chunk_offset + i];- b_vals[i * 2] = dequant_fp4_e2m1(b_val, 0) * scale_b_1;- b_vals[i * 2 + 1] = dequant_fp4_e2m1(b_val, 1) * scale_b_1;- }-- #pragma unroll- for (int r = 0; r < ROWS_PER_WARP; ++r) {- if (!row_valid[r]) continue;-- const int m = base_m + r;- const int64_t a_offset = m * a_s0 + (k_byte_base + chunk_offset) * a_s1 + l * a_s2;-- uint4 a_vec = __ldg(reinterpret_cast<const uint4*>(&a[a_offset]));- unsigned char a_bytes[CHUNK_SIZE];- *reinterpret_cast<uint4*>(&a_bytes[0]) = a_vec;-- const int kk_0 = block_0 / 4, kk4_0 = block_0 % 4;- const int kk_1 = block_1 / 4, kk4_1 = block_1 % 4;-- const float scale_a_0 = dequant_fp8_e4m3(__ldg(&sfa[sfa_m_parts[r] + kk4_0 * sfa_s3 + kk_0 * sfa_s4]));- const float scale_a_1 = dequant_fp8_e4m3(__ldg(&sfa[sfa_m_parts[r] + kk4_1 * sfa_s3 + kk_1 * sfa_s4]));-- #pragma unroll- for (int i = 0; i < 8; ++i) {- const unsigned char a_val = a_bytes[i];- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a_0, b_vals[i * 2], thread_acc[r]);- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a_0, b_vals[i * 2 + 1], thread_acc[r]);- }- #pragma unroll- for (int i = 8; i < CHUNK_SIZE; ++i) {- const unsigned char a_val = a_bytes[i];- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a_1, b_vals[i * 2], thread_acc[r]);- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a_1, b_vals[i * 2 + 1], thread_acc[r]);- }- }- }- }-- 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;+ // Load 8 bytes = 16 FP4 elements, but we only need 8 for one scale group+ // Actually, let's load 4 bytes = 8 FP4 elements = half a scale group+ // Wait - 16 elements per scale group, 8 bytes per scale group+ // Let's load 8 bytes with uint2, process 16 elements- #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, 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]));- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 0) * scale_a, b_fp4_0, thread_acc[r]);- thread_acc[r] = __fmaf_rn(dequant_fp4_e2m1(a_val, 1) * scale_a, b_fp4_1, thread_acc[r]);- }- }+ int k_byte_start = k_scale * 8;+ uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte_start);+ uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte_start);++ // Convert first 4 bytes (8 FP4 values) of A+ uint32_t a_f16_0, a_f16_1, a_f16_2, a_f16_3;+ cvt_f4x8_to_f16x8(a_vec.x, a_f16_0, a_f16_1, a_f16_2, a_f16_3);++ // Convert second 4 bytes (8 FP4 values) of A+ uint32_t a_f16_4, a_f16_5, a_f16_6, a_f16_7;+ cvt_f4x8_to_f16x8(a_vec.y, a_f16_4, a_f16_5, a_f16_6, a_f16_7);++ // Convert first 4 bytes (8 FP4 values) of B+ uint32_t b_f16_0, b_f16_1, b_f16_2, b_f16_3;+ cvt_f4x8_to_f16x8(b_vec.x, b_f16_0, b_f16_1, b_f16_2, b_f16_3);++ // Convert second 4 bytes (8 FP4 values) of B+ uint32_t b_f16_4, b_f16_5, b_f16_6, b_f16_7;+ cvt_f4x8_to_f16x8(b_vec.y, b_f16_4, b_f16_5, b_f16_6, b_f16_7);++ // Compute dot product using half2 operations+ half2* a_h2_0 = reinterpret_cast<half2*>(&a_f16_0);+ half2* a_h2_1 = reinterpret_cast<half2*>(&a_f16_1);+ half2* a_h2_2 = reinterpret_cast<half2*>(&a_f16_2);+ half2* a_h2_3 = reinterpret_cast<half2*>(&a_f16_3);+ half2* a_h2_4 = reinterpret_cast<half2*>(&a_f16_4);+ half2* a_h2_5 = reinterpret_cast<half2*>(&a_f16_5);+ half2* a_h2_6 = reinterpret_cast<half2*>(&a_f16_6);+ half2* a_h2_7 = reinterpret_cast<half2*>(&a_f16_7);++ half2* b_h2_0 = reinterpret_cast<half2*>(&b_f16_0);+ half2* b_h2_1 = reinterpret_cast<half2*>(&b_f16_1);+ half2* b_h2_2 = reinterpret_cast<half2*>(&b_f16_2);+ half2* b_h2_3 = reinterpret_cast<half2*>(&b_f16_3);+ half2* b_h2_4 = reinterpret_cast<half2*>(&b_f16_4);+ half2* b_h2_5 = reinterpret_cast<half2*>(&b_f16_5);+ half2* b_h2_6 = reinterpret_cast<half2*>(&b_f16_6);+ half2* b_h2_7 = reinterpret_cast<half2*>(&b_f16_7);++ // Multiply and accumulate+ half2 prod0 = __hmul2(*a_h2_0, *b_h2_0);+ half2 prod1 = __hmul2(*a_h2_1, *b_h2_1);+ half2 prod2 = __hmul2(*a_h2_2, *b_h2_2);+ half2 prod3 = __hmul2(*a_h2_3, *b_h2_3);+ half2 prod4 = __hmul2(*a_h2_4, *b_h2_4);+ half2 prod5 = __hmul2(*a_h2_5, *b_h2_5);+ half2 prod6 = __hmul2(*a_h2_6, *b_h2_6);+ half2 prod7 = __hmul2(*a_h2_7, *b_h2_7);++ // Sum all products+ half2 sum01 = __hadd2(prod0, prod1);+ half2 sum23 = __hadd2(prod2, prod3);+ half2 sum45 = __hadd2(prod4, prod5);+ half2 sum67 = __hadd2(prod6, prod7);++ half2 sum0123 = __hadd2(sum01, sum23);+ half2 sum4567 = __hadd2(sum45, sum67);++ half2 sum_all = __hadd2(sum0123, sum4567);++ float local_sum = __half2float(sum_all.x) + __half2float(sum_all.y);++ acc += local_sum * scale;}- #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) {- c[m * c_s0 + l * c_s2] = __float2half(sum);- }- }+ acc = warp_reduce_sum(acc);+ if (lane_id == 0) {+ C[row * c_s0 + batch * c_s2] = __float2half(acc);}}- torch::Tensor nvfp4_gemv(- torch::Tensor a,- torch::Tensor b,- torch::Tensor sfa_permuted,- torch::Tensor sfb_permuted,+ torch::Tensor gemv_cuda(+ torch::Tensor a, torch::Tensor b,+ torch::Tensor sfa, torch::Tensor sfb,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");+ int M = a.size(0), K = a.size(1) * 2, L = a.size(2);- 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;- }- cudaMemcpyToSymbol(fp8_e4m3_lut, host_fp8_lut, 256 * sizeof(float));- fp8_lut_initialized = true;- }+ const int WARPS_PER_BLOCK = 8;+ dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);+ dim3 block(32 * WARPS_PER_BLOCK);- 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);- size_t smem_size = K_bytes + 4 + K / 16 + 1;-- // Adaptive configuration: favor parallelism for high-L cases- int total_work = M * L;- bool high_batch = (L >= 4);-- if (high_batch) {- // Use ROWS_PER_WARP=2, VEC_SIZE=32 for better parallelism- constexpr int ROWS_PER_WARP = 2;- constexpr int VEC_SIZE = 32;- const int rows_per_block = block.y * ROWS_PER_WARP;- dim3 grid((M + rows_per_block - 1) / rows_per_block, L);-- nvfp4_gemv_kernel<ROWS_PER_WARP, VEC_SIZE><<<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)- );- } else {- // Use ROWS_PER_WARP=4, VEC_SIZE=32 for better arithmetic intensity- constexpr int ROWS_PER_WARP = 4;- constexpr int VEC_SIZE = 32;- const int rows_per_block = block.y * ROWS_PER_WARP;- dim3 grid((M + rows_per_block - 1) / rows_per_block, L);-- nvfp4_gemv_kernel<ROWS_PER_WARP, VEC_SIZE><<<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));-+ gemv_v5<<<grid, block>>>(+ reinterpret_cast<const uint8_t*>(a.data_ptr()),+ reinterpret_cast<const uint8_t*>(b.data_ptr()),+ reinterpret_cast<const uint8_t*>(sfa.data_ptr()),+ reinterpret_cast<const uint8_t*>(sfb.data_ptr()),+ reinterpret_cast<half*>(c.data_ptr()),+ M, K, L,+ a.stride(0), a.stride(2),+ b.stride(2),+ sfa.stride(0), sfa.stride(1), sfa.stride(2), sfa.stride(3), sfa.stride(4), sfa.stride(5),+ sfb.stride(3), sfb.stride(4), sfb.stride(5),+ c.stride(0), c.stride(2)+ );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- );+ cpp_source = """+ torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);"""- nvfp4_module = load_inline(- name='nvfp4_gemv_p0_optimized',- 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']+ module = load_inline(+ name='gemv_v5',+ cpp_sources=cpp_source,+ cuda_sources=cuda_source,+ functions=['gemv_cuda'],+ extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17', '--generate-code=arch=compute_100a,code=sm_100a'],+ verbose=True)- def custom_kernel(data):- """- Adaptive P0 Optimized NVFP4 GEMV:- - Adaptive ROWS_PER_WARP: 2 for L>=4 (parallelism), 4 for L<4 (intensity)- - VEC_SIZE=32 with chunked processing to reduce register pressure- - __ldg() for cache-optimized loads- - Bank conflict padding-- Expected: 15-25% improvement across all cases- """- a, b, _, _, sfa_permuted, sfb_permuted, c = data- return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)No newline at end of file+ def custom_kernel(data: input_t) -> output_t:+ a, b, sfa, sfb, sfa_perm, sfb_perm, c = data+ return module.gemv_cuda(a, b, sfa_perm, sfb_perm, c)No newline at end of file
scrolls · 501 diff lines total
Best evidence level for this revision: reported
JSON