submission 89542
amackenzie-jumptrading · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 401 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-89542?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:8e36e59f9ba5d8f69ca6b4be91667db1e0d363656dfe47eed006f4a788c7d737
license declaredunknown
license concludedunknown
authorsamackenzie-jumptrading
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {shared-memory
__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {vector-width = half2
half2* s_lut_pair = (half2*)smem;Kernel source
submission.py401 lines
import torch
from torch.utils.cpp_extension import load_inline
# -----------------------------------------------------------------------------
# CUDA Kernel Source
# -----------------------------------------------------------------------------
cuda_source = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#define WARP_SIZE 32
// Small LUT for initialization
__device__ __constant__ float C_FP4_LUT_CONST[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 decode_fp8_fast(uint8_t x) {
if (x == 0x80) return 0.0f;
uint32_t val = x;
uint32_t sign = (val & 0x80) << 24;
uint32_t exp = (val & 0x78) >> 3;
uint32_t mant = (val & 0x07);
uint32_t exp32 = exp + 120;
return (val == 0) ? 0.0f : __int_as_float(sign | (exp32 << 23) | (mant << 20));
}
__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {
uint32_t smem = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], 16;" :: "r"(smem), "l"(glob_ptr));
}
__device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void cp_async_wait_all() { asm volatile("cp.async.wait_group 0;"); }
template <int WARPS_PER_BLOCK>
__global__ void __launch_bounds__(WARPS_PER_BLOCK * WARP_SIZE) fp4_gemv_b200_sol(
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_bytes, int L,
int stride_am, int stride_ak, int stride_al,
int stride_bm, int stride_bk, int stride_bl,
int stride_sfam, int stride_sfak, int stride_sfal,
int stride_sfbm, int stride_sfbk, int stride_sfbl,
int stride_cm, int stride_cl
) {
extern __shared__ char smem[];
// 1. LUT in Shared Memory (256 * 4 bytes = 1KB)
// Maps byte (2x FP4) -> half2 (2x half decoded)
// Optimization: Reduced size from float2 (8B) to half2 (4B)
half2* s_lut_pair = (half2*)smem;
// 2. Decoded B in Shared Memory
// Offset: 2048 bytes (aligned safe margin)
half* s_b_decoded = (half*)(smem + 2048);
int k_elements = K_bytes * 2;
int num_chunks = (k_elements + 31) / 32;
int decoded_size_bytes = num_chunks * 72;
decoded_size_bytes = (decoded_size_bytes + 15) & ~15;
uint8_t* s_b_raw = (uint8_t*)((char*)s_b_decoded + decoded_size_bytes);
uint8_t* s_sfb_raw = (uint8_t*)(s_b_raw + K_bytes);
int tid = threadIdx.x + threadIdx.y * blockDim.x;
// --- Init Huge LUT (first 256 threads) ---
if (tid < 256) {
int lo = tid & 0x0F;
int hi = (tid >> 4) & 0x0F;
// Convert to half2 directly during init
s_lut_pair[tid] = __float22half2_rn(make_float2(C_FP4_LUT_CONST[lo], C_FP4_LUT_CONST[hi]));
}
int l_idx = blockIdx.y;
const uint8_t* B_g = B + l_idx * stride_bl;
const uint8_t* SFB_g = SFB + l_idx * stride_sfbl;
int num_threads = WARPS_PER_BLOCK * WARP_SIZE;
// --- Async Load B & SFB ---
for (int i = tid * 16; i < K_bytes; i += num_threads * 16) {
if (i + 16 <= K_bytes) cp_async_16(&s_b_raw[i], &B_g[i]);
else for(int j=0; j<16 && i+j<K_bytes; ++j) s_b_raw[i+j] = B_g[i+j];
}
int sfb_bytes = K_bytes / 8;
for (int i = tid * 16; i < sfb_bytes; i += num_threads * 16) {
if (i + 16 <= sfb_bytes) cp_async_16(&s_sfb_raw[i], &SFB_g[i]);
else for(int j=0; j<16 && i+j<sfb_bytes; ++j) s_sfb_raw[i+j] = SFB_g[i+j];
}
cp_async_commit();
cp_async_wait_all();
__syncthreads();
// --- Decode B + SFB -> Padded SMEM ---
for (int k = tid * 32; k < k_elements; k += num_threads * 32) {
if (k >= k_elements) break;
uint4 b_vec = *reinterpret_cast<uint4*>(&s_b_raw[k / 2]);
uint8_t* b_bytes = (uint8_t*)&b_vec;
int sfb_idx = k / 16;
float s0 = decode_fp8_fast(s_sfb_raw[sfb_idx]);
float s1 = decode_fp8_fast(s_sfb_raw[sfb_idx+1]);
int write_base = (k >> 5) * 36;
#pragma unroll
for (int i = 0; i < 32; ++i) {
uint8_t packed = b_bytes[i/2];
int nib = (i & 1) ? (packed >> 4) : (packed & 0xF);
float val = C_FP4_LUT_CONST[nib];
float scale = (i < 16) ? s0 : s1;
s_b_decoded[write_base + i] = __float2half(val * scale);
}
}
__syncthreads();
// --- Math Phase ---
int warp_id = threadIdx.y;
int lane_id = threadIdx.x;
int m = blockIdx.x * WARPS_PER_BLOCK + warp_id;
if (m < M) {
const uint8_t* A_ptr = A + m * stride_am + l_idx * stride_al;
const uint8_t* SFA_ptr = SFA + m * stride_sfam + l_idx * stride_sfal;
float acc = 0.0f;
uint4 a_reg;
ushort sfa_reg;
// Prologue
int k_first = lane_id * 32;
if (k_first < k_elements) {
a_reg = *reinterpret_cast<const uint4*>(A_ptr + k_first / 2);
sfa_reg = *reinterpret_cast<const ushort*>(SFA_ptr + k_first / 16);
}
// Main Loop
for (int k_base = 0; k_base < k_elements; k_base += 1024) {
int k = k_base + lane_id * 32;
// 1. Prefetch
int k_next = k + 1024;
uint4 a_next;
ushort sfa_next;
bool active_next = (k_next < k_elements);
if (active_next) {
a_next = *reinterpret_cast<const uint4*>(A_ptr + k_next / 2);
sfa_next = *reinterpret_cast<const ushort*>(SFA_ptr + k_next / 16);
}
// 2. Compute
if (k < k_elements) {
float sfa0 = decode_fp8_fast(sfa_reg & 0xFF);
float sfa1 = decode_fp8_fast(sfa_reg >> 8);
uint8_t* a_bytes = (uint8_t*)&a_reg;
int read_base_idx = (k >> 5) * 36;
int2* b_vec_ptr = reinterpret_cast<int2*>(&s_b_decoded[read_base_idx]);
int2 b_vals[8];
#pragma unroll
for(int v=0; v<8; ++v) b_vals[v] = b_vec_ptr[v];
half2* b_h2_ptr = (half2*)b_vals;
// Accumulators in half2
half2 sum0_h2 = __float2half2_rn(0.0f);
half2 sum1_h2 = __float2half2_rn(0.0f);
// Loop 0-7
#pragma unroll
for (int i = 0; i < 8; ++i) {
// SMEM Load half2 (4 bytes)
half2 va = s_lut_pair[a_bytes[i]];
half2 vb = b_h2_ptr[i];
// Vectorized FMA
sum0_h2 = __hfma2(va, vb, sum0_h2);
}
// Loop 8-15
#pragma unroll
for (int i = 8; i < 16; ++i) {
half2 va = s_lut_pair[a_bytes[i]];
half2 vb = b_h2_ptr[i];
sum1_h2 = __hfma2(va, vb, sum1_h2);
}
// Reduction: sum .x and .y components and apply scale
// We cast to float here to maintain precision for the accumulation across K
float partial0 = __half2float(sum0_h2.x) + __half2float(sum0_h2.y);
float partial1 = __half2float(sum1_h2.x) + __half2float(sum1_h2.y);
acc += partial0 * sfa0 + partial1 * sfa1;
}
// 3. Update
if (active_next) {
a_reg = a_next;
sfa_reg = sfa_next;
}
}
for (int offset = 16; offset > 0; offset /= 2)
acc += __shfl_down_sync(0xffffffff, acc, offset);
if (lane_id == 0) {
C[m * stride_cm + l_idx * stride_cl] = __float2half(acc);
}
}
}
// Helper template to launch kernel with specific WARPS_PER_BLOCK
template<int WARPS_PER_BLOCK>
void launch_kernel_with_config(
const uint8_t* A, const uint8_t* B, const uint8_t* SFA, const uint8_t* SFB, half* C,
int M, int K_bytes, int L,
int stride_am, int stride_ak, int stride_al,
int stride_bm, int stride_bk, int stride_bl,
int stride_sfam, int stride_sfak, int stride_sfal,
int stride_sfbm, int stride_sfbk, int stride_sfbl,
int stride_cm, int stride_cl,
size_t smem_bytes
) {
dim3 block(32, WARPS_PER_BLOCK, 1);
dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L, 1);
fp4_gemv_b200_sol<WARPS_PER_BLOCK><<<grid, block, smem_bytes>>>(
A, B, SFA, SFB, C,
M, K_bytes, L,
stride_am, stride_ak, stride_al,
stride_bm, stride_bk, stride_bl,
stride_sfam, stride_sfak, stride_sfal,
stride_sfbm, stride_sfbk, stride_sfbl,
stride_cm, stride_cl
);
}
void fp4_gemv_launch(
torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
int stride_am, int stride_ak, int stride_al,
int stride_bm, int stride_bk, int stride_bl,
int stride_sfam, int stride_sfak, int stride_sfal,
int stride_sfbm, int stride_sfbk, int stride_sfbl,
int stride_cm, int stride_cl,
int M, int K_bytes, int L
) {
int K_elements = K_bytes * 2;
int num_chunks = (K_elements + 31) / 32;
size_t size_decoded = num_chunks * 72;
size_decoded = (size_decoded + 15) & ~15;
size_t size_raw = K_bytes + (K_bytes/8);
// SMEM: 2048 (LUT) + Decoded B + Raw + Padding
size_t smem_bytes = 2048 + size_decoded + size_raw + 256;
const uint8_t* A_ptr = A.data_ptr<uint8_t>();
const uint8_t* B_ptr = B.data_ptr<uint8_t>();
const uint8_t* SFA_ptr = SFA.data_ptr<uint8_t>();
const uint8_t* SFB_ptr = SFB.data_ptr<uint8_t>();
half* C_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
// Dispatch based on shape characteristics
// Shape 1: M=7168, K=16384, L=1 -> warps=18
// Shape 2: M=4096, K=7168, L=8 -> warps=16
// Shape 3: M=7168, K=2048, L=4 -> warps=17
// Launch appropriate kernel
if (L == 1) {
launch_kernel_with_config<18>(
A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr,
M, K_bytes, L,
stride_am, stride_ak, stride_al,
stride_bm, stride_bk, stride_bl,
stride_sfam, stride_sfak, stride_sfal,
stride_sfbm, stride_sfbk, stride_sfbl,
stride_cm, stride_cl,
smem_bytes
);
} else {
launch_kernel_with_config<16>(
A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr,
M, K_bytes, L,
stride_am, stride_ak, stride_al,
stride_bm, stride_bk, stride_bl,
stride_sfam, stride_sfak, stride_sfal,
stride_sfbm, stride_sfbk, stride_sfbl,
stride_cm, stride_cl,
smem_bytes
);
}
}
"""
cpp_source = r"""
void fp4_gemv_launch(
torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
int stride_am, int stride_ak, int stride_al,
int stride_bm, int stride_bk, int stride_bl,
int stride_sfam, int stride_sfak, int stride_sfal,
int stride_sfbm, int stride_sfbk, int stride_sfbl,
int stride_cm, int stride_cl,
int M, int K_bytes, int L
);
void fp4_gemv_opt(
torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
int stride_am, int stride_ak, int stride_al,
int stride_bm, int stride_bk, int stride_bl,
int stride_sfam, int stride_sfak, int stride_sfal,
int stride_sfbm, int stride_sfbk, int stride_sfbl,
int stride_cm, int stride_cl,
int M, int K_bytes, int L
) {
fp4_gemv_launch(A, B, SFA, SFB, C,
stride_am, stride_ak, stride_al,
stride_bm, stride_bk, stride_bl,
stride_sfam, stride_sfak, stride_sfal,
stride_sfbm, stride_sfbk, stride_sfbl,
stride_cm, stride_cl,
M, K_bytes, L);
}
"""
try:
from task import input_t, output_t
except ImportError:
import collections
input_t = collections.namedtuple("input_t", ["a", "b", "sfa", "sfb", "sfa_perm", "sfb_perm", "c"])
output_t = torch.Tensor
_cuda_module = None
def get_cuda_module():
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
name="fp4_gemv_b200_sol_v204",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["fp4_gemv_opt"],
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
"-std=c++17",
"-maxrregcount=255",
"--generate-line-info"
],
)
return _cuda_module
def _as_uint8(t: torch.Tensor) -> torch.Tensor:
if t.dtype == torch.uint8: return t
return t.view(torch.uint8)
def _check_device(t):
if not t.is_cuda: raise RuntimeError("All inputs must be on CUDA device")
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, _, _, c = data
_check_device(a); _check_device(b); _check_device(sfa); _check_device(sfb); _check_device(c)
M, K_bytes, L = a.shape
A_u8 = _as_uint8(a)
B_u8 = _as_uint8(b)
SFA_u8 = _as_uint8(sfa)
SFB_u8 = _as_uint8(sfb)
C_view = c.view(M, L)
sam, sak, sal = map(int, A_u8.stride())
sbm, sbk, sbl = map(int, B_u8.stride())
ssfam, ssfak, ssfal = map(int, SFA_u8.stride())
ssfbm, ssfbk, ssfbl = map(int, SFB_u8.stride())
scm, scl = map(int, C_view.stride())
mod = get_cuda_module()
mod.fp4_gemv_opt(
A_u8, B_u8, SFA_u8, SFB_u8, C_view,
sam, sak, sal,
sbm, sbk, sbl,
ssfam, ssfak, ssfal,
ssfbm, ssfbk, ssfbl,
scm, scl,
M, K_bytes, L
)
return c
scrolls · 401 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON