submission 109158
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 303 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109158?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:048a494a47fcb24b51a12cfef6b4390ad155599c6645d0cb45fd38ca9f65bead
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))fp8
__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);shared-memory
__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][M_TILE][BYTES_PER_TILE];vector-width = float2
float2 f0 = __half22float2(acc0);Kernel source
submission.py303 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#define BLOCK_SIZE 128 // 4 warps (vs 32 in original)
#define K_TILE 2048
#define SCALES_PER_TILE (K_TILE / 16) // 160
#define BYTES_PER_TILE (K_TILE / 2) // 1280
#define M_TILE 4 // 4 rows per block (vs 1 in original)
#define NUM_BUFFERS 2
// Async copy macros
#define ASYNC_COPY_16(dst, src) \
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
#define ASYNC_COPY_4(dst, src) \
asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))
#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
#define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")
__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
__half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
static_cast<__nv_fp4x2_storage_t>(byte), __NV_E2M1);
return *reinterpret_cast<__half2*>(&raw);
}
__device__ __forceinline__ float decode_fp8(int8_t byte) {
__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
__half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
return __half2float(__ushort_as_half(raw.x));
}
__device__ __forceinline__ __half2 dot_scaled_4bytes(uint32_t a4, uint32_t b4, __half2 scale_h2) {
uint32_t b0, b1, b2, b3, a0, a1, a2, a3;
asm("bfe.u32 %0, %1, 0, 8;" : "=r"(b0) : "r"(b4));
asm("bfe.u32 %0, %1, 8, 8;" : "=r"(b1) : "r"(b4));
asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b2) : "r"(b4));
asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b3) : "r"(b4));
asm("bfe.u32 %0, %1, 0, 8;" : "=r"(a0) : "r"(a4));
asm("bfe.u32 %0, %1, 8, 8;" : "=r"(a1) : "r"(a4));
asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a2) : "r"(a4));
asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a3) : "r"(a4));
__half2 acc = __hmul2(decode_fp4x2(a0), __hmul2(decode_fp4x2(b0), scale_h2));
acc = __hfma2(decode_fp4x2(a1), __hmul2(decode_fp4x2(b1), scale_h2), acc);
acc = __hfma2(decode_fp4x2(a2), __hmul2(decode_fp4x2(b2), scale_h2), acc);
acc = __hfma2(decode_fp4x2(a3), __hmul2(decode_fp4x2(b3), scale_h2), acc);
return acc;
}
__device__ __forceinline__ float warp_reduce_sum(float val) {
val += __shfl_down_sync(0xffffffff, val, 16);
val += __shfl_down_sync(0xffffffff, val, 8);
val += __shfl_down_sync(0xffffffff, val, 4);
val += __shfl_down_sync(0xffffffff, val, 2);
val += __shfl_down_sync(0xffffffff, val, 1);
return val;
}
__global__ __launch_bounds__(BLOCK_SIZE)
void gemv_mtiled_kernel(
const int8_t* __restrict__ a,
const int8_t* __restrict__ b,
const int8_t* __restrict__ sfa,
const int8_t* __restrict__ sfb,
half* __restrict__ c,
int M, int K, int L, int N_rows
) {
// Shared memory: 4 A rows + 1 B (shared across all 4 warps!)
__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][M_TILE][BYTES_PER_TILE];
__shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][M_TILE][SCALES_PER_TILE];
__shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
__shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];
// Block handles M_TILE consecutive rows: [m_base, m_base+1, m_base+2, m_base+3]
const int m_base = blockIdx.x * M_TILE;
const int batch_id = blockIdx.y;
const int tid = threadIdx.x;
const int warp_id = tid >> 5;
const int lane = tid & 31;
// Bounds
const int valid_rows = min(M_TILE, M - m_base);
const bool my_row_valid = (warp_id < valid_rows);
const int my_m = m_base + warp_id;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const int tile_count = K_bytes / BYTES_PER_TILE;
const int remainder_sf_start = (tile_count * BYTES_PER_TILE) / 8;
const bool has_remainder = (remainder_sf_start < K_sf);
const int remainder_scales = K_sf - remainder_sf_start;
const size_t a_batch_stride = (size_t)M * K_bytes;
const size_t b_batch_stride = (size_t)N_rows * K_bytes;
const size_t sfa_batch_stride = (size_t)M * K_sf;
const size_t sfb_batch_stride = (size_t)N_rows * K_sf;
const uint8_t* batch_a = reinterpret_cast<const uint8_t*>(a) + batch_id * a_batch_stride;
const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + batch_id * b_batch_stride;
const uint8_t* batch_sfa = reinterpret_cast<const uint8_t*>(sfa) + batch_id * sfa_batch_stride;
const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + batch_id * sfb_batch_stride;
float acc = 0.0f;
int buf = 0;
auto issue_tile_async = [&](int b_idx, int tile) {
const int base_byte = tile * BYTES_PER_TILE;
const int base_sf = tile * SCALES_PER_TILE;
const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
}
if (my_row_valid) {
const uint8_t* row_a = batch_a + my_m * K_bytes;
const uint8_t* row_sfa = batch_sfa + my_m * K_sf;
const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][warp_id][0]);
const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][warp_id][0]);
for (int i = lane * 16; i < BYTES_PER_TILE; i += 32 * 16) {
ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
}
for (int i = lane * 4; i < SCALES_PER_TILE; i += 32 * 4) {
ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
}
}
ASYNC_COMMIT();
};
auto issue_remainder_async = [&](int b_idx) {
const int base_byte = remainder_sf_start << 3;
const int rem_bytes = remainder_scales << 3;
const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
for (int i = tid * 16; i < rem_bytes; i += BLOCK_SIZE * 16) {
ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {
ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + remainder_sf_start + i);
}
if (my_row_valid) {
const uint8_t* row_a = batch_a + my_m * K_bytes;
const uint8_t* row_sfa = batch_sfa + my_m * K_sf;
const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][warp_id][0]);
const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][warp_id][0]);
for (int i = lane * 16; i < rem_bytes; i += 32 * 16) {
ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
}
for (int i = lane * 4; i < remainder_scales; i += 32 * 4) {
ASYNC_COPY_4(sh_sfa_base + i, row_sfa + remainder_sf_start + i);
}
}
ASYNC_COMMIT();
};
// Main K-tile loop with double buffering
if (tile_count > 0) {
issue_tile_async(0, 0);
ASYNC_WAIT_ALL();
__syncthreads();
for (int tile = 0; tile < tile_count; ++tile) {
if (tile + 1 < tile_count) {
issue_tile_async(buf ^ 1, tile + 1);
} else if (has_remainder) {
issue_remainder_async(buf ^ 1);
}
if (my_row_valid) {
float tile_acc = 0.0f;
#pragma unroll 4
for (int sf = lane; sf < SCALES_PER_TILE; sf += 32) {
float scale = decode_fp8(static_cast<int8_t>(sh_sfa[buf][warp_id][sf])) *
decode_fp8(static_cast<int8_t>(sh_sfb[buf][sf]));
__half2 scale_h2 = __half2half2(__float2half(scale));
int byte_base = sf << 3;
uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base]);
uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base]);
uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base + 4]);
uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base + 4]);
__half2 acc0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
__half2 acc1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
float2 f0 = __half22float2(acc0);
float2 f1 = __half22float2(acc1);
tile_acc += f0.x + f0.y + f1.x + f1.y;
}
acc += tile_acc;
}
if (tile + 1 < tile_count || has_remainder) {
ASYNC_WAIT_ALL();
__syncthreads();
buf ^= 1;
}
}
}
// Remainder
if (has_remainder) {
if (tile_count == 0) {
issue_remainder_async(0);
ASYNC_WAIT_ALL();
__syncthreads();
buf = 0;
}
if (my_row_valid) {
float rem_acc = 0.0f;
for (int sf = lane; sf < remainder_scales; sf += 32) {
float scale = decode_fp8(static_cast<int8_t>(sh_sfa[buf][warp_id][sf])) *
decode_fp8(static_cast<int8_t>(sh_sfb[buf][sf]));
__half2 scale_h2 = __half2half2(__float2half(scale));
int byte_base = sf << 3;
uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base]);
uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base]);
uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base + 4]);
uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base + 4]);
__half2 acc0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
__half2 acc1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
float2 f0 = __half22float2(acc0);
float2 f1 = __half22float2(acc1);
rem_acc += f0.x + f0.y + f1.x + f1.y;
}
acc += rem_acc;
}
}
if (my_row_valid) {
float warp_sum = warp_reduce_sum(acc);
if (lane == 0) {
c[(size_t)my_m + (size_t)batch_id * M] = __float2half(warp_sum);
}
}
}
void run_nvfp4_gemv(
torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
int num_m_blocks = (m + M_TILE - 1) / M_TILE;
dim3 grid(num_m_blocks, l);
dim3 block(BLOCK_SIZE);
gemv_mtiled_kernel<<<grid, block>>>(
reinterpret_cast<const int8_t*>(A.data_ptr()),
reinterpret_cast<const int8_t*>(B.data_ptr()),
reinterpret_cast<const int8_t*>(SFA.data_ptr()),
reinterpret_cast<const int8_t*>(SFB.data_ptr()),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
m, k, l, n_pad
);
}
'''
CPP_SRC = r'''
#include <torch/extension.h>
void run_nvfp4_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);
'''
_cuda_module = None
def get_cuda_module():
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
name='nvfp4_gemv_mtiled_v3',
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=['run_nvfp4_gemv'],
extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a'],
verbose=False,
)
return _cuda_module
def custom_kernel(data: input_t) -> output_t:
a, b, sfa_ref, sfb_ref, _, _, c = data
module = get_cuda_module()
m, k_packed, l = a.shape
k = k_packed * 2
n_pad = b.shape[0]
if not sfa_ref.is_cuda:
sfa_ref = sfa_ref.to(a.device)
if not sfb_ref.is_cuda:
sfb_ref = sfb_ref.to(a.device)
c_out = c.squeeze(1)
module.run_nvfp4_gemv(c_out, a.view(torch.uint8), b.view(torch.uint8),
sfa_ref.view(torch.uint8), sfb_ref.view(torch.uint8),
m, k, l, n_pad)
return c
scrolls · 303 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 108012.
- """- NVFP4 GEMV - M-tiled Kernel (Non-Persistent, High Parallelism)-- Key idea: Same as submission_1126_rawcuda.py, but each block processes 4 rows.- - Original: grid(M, L), 1 row per block- - This: grid(M/4, L), 4 rows per block with 4 warps-- Grid: (M/4, L) = (1792, 1) for M=7168- - Each block: 4 warps, each warp computes 1 row- - B loaded once per block, shared by 4 rows (4× reuse)- - Same high parallelism as original!- """-import torchfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t⋯ 4 unchanged lines#include <cuda_fp4.h>#include <cuda_fp8.h>- // ============================================================================- // Configuration- // ============================================================================#define BLOCK_SIZE 128 // 4 warps (vs 32 in original)- #define K_TILE 2560+ #define K_TILE 2048#define SCALES_PER_TILE (K_TILE / 16) // 160#define BYTES_PER_TILE (K_TILE / 2) // 1280#define M_TILE 4 // 4 rows per block (vs 1 in original)⋯ 7 unchanged lines#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")#define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")- // ============================================================================- // FP4/FP8 Conversion- // ============================================================================__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {__half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(static_cast<__nv_fp4x2_storage_t>(byte), __NV_E2M1);⋯ 33 unchanged linesreturn val;}- // ============================================================================- // M-tiled GEMV kernel - 4 warps, each handles one row- // Grid: (M/4, L) - same parallelism as original, with 4x B reuse- // ============================================================================__global__ __launch_bounds__(BLOCK_SIZE)void gemv_mtiled_kernel(const int8_t* __restrict__ a,⋯ 13 unchanged linesconst int m_base = blockIdx.x * M_TILE;const int batch_id = blockIdx.y;const int tid = threadIdx.x;- const int warp_id = tid >> 5; // 0-3: which row this warp handles- const int lane = tid & 31; // 0-31: lane within warp+ const int warp_id = tid >> 5;+ const int lane = tid & 31;// Boundsconst int valid_rows = min(M_TILE, M - m_base);const bool my_row_valid = (warp_id < valid_rows);- const int my_m = m_base + warp_id; // This warp's row+ const int my_m = m_base + warp_id;const int K_bytes = K / 2;const int K_sf = K / 16;⋯ 15 unchanged linesfloat acc = 0.0f;int buf = 0;- // Load function: B shared, each warp loads its own A rowauto issue_tile_async = [&](int b_idx, int tile) {const int base_byte = tile * BYTES_PER_TILE;const int base_sf = tile * SCALES_PER_TILE;- // B: All 128 threads cooperate to load (fast!)const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {⋯ 3 unchanged linesASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);}- // A: Each warp loads its own row IN PARALLELif (my_row_valid) {const uint8_t* row_a = batch_a + my_m * K_bytes;const uint8_t* row_sfa = batch_sfa + my_m * K_sf;⋯ 48 unchanged linesissue_remainder_async(buf ^ 1);}- // Each warp computes its rowif (my_row_valid) {float tile_acc = 0.0f;- #pragma unroll 5+ #pragma unroll 4for (int sf = lane; sf < SCALES_PER_TILE; sf += 32) {float scale = decode_fp8(static_cast<int8_t>(sh_sfa[buf][warp_id][sf])) *decode_fp8(static_cast<int8_t>(sh_sfb[buf][sf]));⋯ 49 unchanged lines}}- // Output: each warp writes its own rowif (my_row_valid) {float warp_sum = warp_reduce_sum(acc);if (lane == 0) {⋯ 2 unchanged lines}}- // ============================================================================- // Host wrapper- // ============================================================================void run_nvfp4_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,torch::Tensor SFA, torch::Tensor SFB,int m, int k, int l, int n_pad) {- // Grid: (M/4, L) - same parallelism as original grid(M, L)!- int num_m_blocks = (m + M_TILE - 1) / M_TILE; // ceil(M/4)+ int num_m_blocks = (m + M_TILE - 1) / M_TILE;- dim3 grid(num_m_blocks, l); // (1792, 1) for M=7168, L=1- dim3 block(BLOCK_SIZE); // 128 threads (4 warps)+ dim3 grid(num_m_blocks, l);+ dim3 block(BLOCK_SIZE);gemv_mtiled_kernel<<<grid, block>>>(reinterpret_cast<const int8_t*>(A.data_ptr()),
scrolls · 131 diff lines total
Best evidence level for this revision: reported
JSON