submission 107168
XoTic · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 453 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-107168?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:32fad00e4950d0bc4f00fb0df952a9a24fde3f2359aadc7e1fcff702772c26bb
license declaredunknown
license concludedunknown
authorsXoTic
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
"cp.async.cg.shared.global [%0], [%1], 16;\n"shared-memory
extern __shared__ half b_shared[];vector-width = uint4
const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[byte_off]);Kernel source
submission.py453 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
cuda_source = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>
__device__ __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__ size_t row_scale_pointer(int m, int rest_m, int rest_k, int L) {
const int block = m >> 7;
const int idx32 = m & 31;
const int idx4 = (m >> 5) & 0x3;
const size_t base = static_cast<size_t>(((idx32 * 4 + idx4) * rest_m + block) * 4);
return base * rest_k * L;
}
__device__ __forceinline__ int get_scale_idx_b(int kb, int rest_k, int L, int l) {
const int kk = kb >> 2;
const int kk4 = kb & 0x3;
return (kk4 * rest_k + kk) * L + l;
}
__global__ void decode_b_kernel(
const uint8_t* __restrict__ b_packed,
const half* __restrict__ sfb_scales,
half* __restrict__ b_decoded,
int K,
int L,
int rest_k
) {
const int K_half = K >> 1;
const int total_bytes = K_half * L;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < total_bytes; idx += blockDim.x * gridDim.x) {
const int byte_in_k = idx / L;
const int l = idx - byte_in_k * L;
const uint8_t packed = b_packed[idx];
const int k = byte_in_k << 1;
const int kb = k >> 4;
const float scale = __half2float(__ldg(&sfb_scales[get_scale_idx_b(kb, rest_k, L, l)]));
const float v0 = fp4_lut[packed & 0xF] * scale;
const float v1 = fp4_lut[packed >> 4] * scale;
const size_t base = static_cast<size_t>(l) * K + k;
b_decoded[base] = __float2half(v0);
b_decoded[base + 1] = __float2half(v1);
}
}
// Double-buffered L=1 kernel with cp.async
// CHUNK_SIZE is the number of half elements per chunk (must be multiple of 8 for cp.async.cg alignment)
constexpr int L1_CHUNK_SIZE = 2048;
__global__ void __launch_bounds__(256) gemv_L1_kernel(
const uint8_t* __restrict__ a_packed,
const half* __restrict__ b_decoded,
const half* __restrict__ sfa_scales,
half* __restrict__ c,
int M, int K,
int rest_m, int rest_k
) {
// Double buffer: two chunks of L1_CHUNK_SIZE halfs each
extern __shared__ half b_shared[];
half* buf0 = b_shared;
half* buf1 = b_shared + L1_CHUNK_SIZE;
const int tid = threadIdx.x;
const int warp_id = tid >> 5;
const int lane_id = tid & 31;
constexpr int ROWS_PER_BLOCK = 8;
const int m = blockIdx.x * ROWS_PER_BLOCK + warp_id;
const int K_half = K >> 1;
const int num_chunks = (K + L1_CHUNK_SIZE - 1) / L1_CHUNK_SIZE;
// Prefetch first chunk into buf0
const int elems_chunk0 = (L1_CHUNK_SIZE < K) ? L1_CHUNK_SIZE : K;
// Use cp.async with 16-byte (8 half) granularity
for (int idx = tid * 8; idx < elems_chunk0; idx += blockDim.x * 8) {
if (idx + 8 <= elems_chunk0) {
asm volatile(
"cp.async.cg.shared.global [%0], [%1], 16;\n"
:
: "r"(static_cast<unsigned>(__cvta_generic_to_shared(&buf0[idx]))),
"l"(&b_decoded[idx])
);
}
}
asm volatile("cp.async.commit_group;\n");
if (m >= M) {
// Still need to participate in async completion
for (int chunk = 0; chunk < num_chunks; ++chunk) {
asm volatile("cp.async.wait_group 0;\n");
__syncthreads();
if (chunk + 1 < num_chunks) {
asm volatile("cp.async.commit_group;\n");
}
__syncthreads();
}
return;
}
const uint8_t* row_ptr = a_packed + static_cast<size_t>(m) * K_half;
const half* row_scales = sfa_scales + row_scale_pointer(m, rest_m, rest_k, 1);
float sum = 0.0f;
for (int chunk = 0; chunk < num_chunks; ++chunk) {
const int k_chunk_start = chunk * L1_CHUNK_SIZE;
const int k_chunk_end = ((k_chunk_start + L1_CHUNK_SIZE) < K) ? (k_chunk_start + L1_CHUNK_SIZE) : K;
const int chunk_elems = k_chunk_end - k_chunk_start;
// Prefetch next chunk into alternate buffer
if (chunk + 1 < num_chunks) {
const int next_k_start = (chunk + 1) * L1_CHUNK_SIZE;
const int next_k_end = ((next_k_start + L1_CHUNK_SIZE) < K) ? (next_k_start + L1_CHUNK_SIZE) : K;
const int next_elems = next_k_end - next_k_start;
half* next_buf = (chunk & 1) ? buf0 : buf1;
for (int idx = tid * 8; idx < next_elems; idx += blockDim.x * 8) {
if (idx + 8 <= next_elems) {
asm volatile(
"cp.async.cg.shared.global [%0], [%1], 16;\n"
:
: "r"(static_cast<unsigned>(__cvta_generic_to_shared(&next_buf[idx]))),
"l"(&b_decoded[next_k_start + idx])
);
}
}
asm volatile("cp.async.commit_group;\n");
}
// Wait for current chunk
asm volatile("cp.async.wait_group 1;\n");
__syncthreads();
// Compute on current chunk
half* cur_buf = (chunk & 1) ? buf1 : buf0;
const int byte_start = k_chunk_start >> 1;
const int byte_end = k_chunk_end >> 1;
for (int byte_off = byte_start + (lane_id << 4); byte_off < byte_end; byte_off += 512) {
const int bytes_to_read = ((byte_off + 16) <= byte_end) ? 16 : (byte_end - byte_off);
if (bytes_to_read < 16) break; // Skip partial reads for simplicity
const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[byte_off]);
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
const int k_start = byte_off << 1;
const int kb0 = k_start >> 4;
const int kb1 = kb0 + 1;
const int kk0 = kb0 >> 2;
const int kk4_0 = kb0 & 0x3;
const int kk1 = kb1 >> 2;
const int kk4_1 = kb1 & 0x3;
const float sa0 = __half2float(__ldg(&row_scales[kk4_0 * rest_k + kk0]));
const float sa1 = __half2float(__ldg(&row_scales[kk4_1 * rest_k + kk1]));
#pragma unroll
for (int i = 0; i < 8; ++i) {
const uint8_t ab = a_bytes[i];
const int k = k_start + (i << 1);
const int k_local = k - k_chunk_start;
const float a0 = fp4_lut[ab & 0xF] * sa0;
const float a1 = fp4_lut[ab >> 4] * sa0;
const float b0 = __half2float(cur_buf[k_local]);
const float b1 = __half2float(cur_buf[k_local + 1]);
sum = __fmaf_rn(a0, b0, sum);
sum = __fmaf_rn(a1, b1, sum);
}
#pragma unroll
for (int i = 8; i < 16; ++i) {
const uint8_t ab = a_bytes[i];
const int k = k_start + (i << 1);
const int k_local = k - k_chunk_start;
const float a0 = fp4_lut[ab & 0xF] * sa1;
const float a1 = fp4_lut[ab >> 4] * sa1;
const float b0 = __half2float(cur_buf[k_local]);
const float b1 = __half2float(cur_buf[k_local + 1]);
sum = __fmaf_rn(a0, b0, sum);
sum = __fmaf_rn(a1, b1, sum);
}
}
__syncthreads();
}
// Final wait for any remaining async ops
asm volatile("cp.async.wait_group 0;\n");
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
if (lane_id == 0) {
c[m] = __float2half(sum);
}
}
template<int L, int ROWS_PER_BLOCK, int L_PER_WARP>
__global__ void __launch_bounds__(256) gemv_multi_kernel(
const uint8_t* __restrict__ a_packed,
const half* __restrict__ b_decoded,
const half* __restrict__ sfa_scales,
half* __restrict__ c,
int M, int K,
int rest_m, int rest_k
) {
constexpr int CHUNKS_PER_ROW = L / L_PER_WARP;
constexpr int KHALF_PER_VEC = 16 / L;
const int tid = threadIdx.x;
const int warp_id = tid >> 5;
const int lane_id = tid & 31;
const int row_in_block = warp_id / CHUNKS_PER_ROW;
if (row_in_block >= ROWS_PER_BLOCK) return;
const int chunk_id = warp_id % CHUNKS_PER_ROW;
const int l_base = chunk_id * L_PER_WARP;
const int m = blockIdx.x * ROWS_PER_BLOCK + row_in_block;
if (m >= M) return;
const int K_half = K >> 1;
const int iter_bound = K_half / KHALF_PER_VEC;
const size_t a_row_base = static_cast<size_t>(m) * K_half * L;
const uint8_t* row_ptr = a_packed + a_row_base;
const half* row_scales = sfa_scales + row_scale_pointer(m, rest_m, rest_k, L);
const half* b_ptrs[L_PER_WARP];
#pragma unroll
for (int lp = 0; lp < L_PER_WARP; ++lp) {
const int l = l_base + lp;
b_ptrs[lp] = (l < L) ? (b_decoded + static_cast<size_t>(l) * K) : nullptr;
}
float sum[L_PER_WARP] = {0.0f};
for (int vec_idx = lane_id; vec_idx < iter_bound; vec_idx += 32) {
const int k_half_base = vec_idx * KHALF_PER_VEC;
const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[k_half_base * L]);
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
const int kb = (k_half_base << 1) >> 4;
const int kk = kb >> 2;
const int kk4 = kb & 0x3;
float sa_vals[L_PER_WARP];
#pragma unroll
for (int lp = 0; lp < L_PER_WARP; ++lp) {
const int l = l_base + lp;
if (l < L) {
const int scale_idx = (kk4 * rest_k + kk) * L + l;
sa_vals[lp] = __half2float(__ldg(&row_scales[scale_idx]));
} else {
sa_vals[lp] = 0.0f;
}
}
#pragma unroll
for (int i = 0; i < KHALF_PER_VEC; ++i) {
const int byte_off = i * L + l_base;
const int k_pair = ((k_half_base + i) << 1);
#pragma unroll
for (int lp = 0; lp < L_PER_WARP; ++lp) {
const int l = l_base + lp;
if (l >= L) continue;
const uint8_t ab = a_bytes[byte_off + lp];
const float a0 = fp4_lut[ab & 0xF] * sa_vals[lp];
const float a1 = fp4_lut[ab >> 4] * sa_vals[lp];
const half* b_vec = b_ptrs[lp];
const float b0 = __half2float(b_vec[k_pair]);
const float b1 = __half2float(b_vec[k_pair + 1]);
sum[lp] = __fmaf_rn(a0, b0, sum[lp]);
sum[lp] = __fmaf_rn(a1, b1, sum[lp]);
}
}
}
#pragma unroll
for (int lp = 0; lp < L_PER_WARP; ++lp) {
const int l = l_base + lp;
if (l >= L) continue;
float val = sum[lp];
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
if (lane_id == 0) {
c[m * L + l] = __float2half(val);
}
}
}
extern "C" void decode_nvfp4_vector(
const uint8_t* b_packed,
const half* sfb_scales,
half* b_decoded,
int K, int L,
int rest_k
) {
const int threads = 256;
const int total_pairs = (K >> 1) * L;
int blocks = (total_pairs + threads - 1) / threads;
blocks = blocks == 0 ? 1 : blocks;
blocks = blocks > 65535 ? 65535 : blocks;
decode_b_kernel<<<blocks, threads>>>(b_packed, sfb_scales, b_decoded, K, L, rest_k);
}
extern "C" void launch_nvfp4_kernel(
const uint8_t* a_packed,
const half* b_decoded,
const half* sfa_scales,
half* c,
int M, int K, int L,
int rest_m, int rest_k
) {
const int threads = 256;
if (L == 1) {
constexpr int ROWS_PER_BLOCK = 8;
constexpr int L1_CHUNK_SIZE = 2048;
const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
// Double buffer: 2 chunks of L1_CHUNK_SIZE halfs
const int smem = 2 * L1_CHUNK_SIZE * sizeof(half);
gemv_L1_kernel<<<blocks, threads, smem>>>(
a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
);
} else if (L == 2) {
constexpr int ROWS_PER_BLOCK = 4;
const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
gemv_multi_kernel<2, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
);
} else if (L == 4) {
constexpr int ROWS_PER_BLOCK = 4;
const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
gemv_multi_kernel<4, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
);
} else if (L == 8) {
constexpr int ROWS_PER_BLOCK = 2;
const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
gemv_multi_kernel<8, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
);
}
}
"""
cpp_source = """
#include <torch/extension.h>
extern "C" void decode_nvfp4_vector(
const uint8_t* b_packed,
const at::Half* sfb_scales,
at::Half* b_decoded,
int K, int L,
int rest_k
);
extern "C" void launch_nvfp4_kernel(
const uint8_t* a_packed,
const at::Half* b_decoded,
const at::Half* sfa_scales,
at::Half* c,
int M, int K, int L,
int rest_m, int rest_k
);
torch::Tensor nvfp4_gemv(
torch::Tensor a_packed,
torch::Tensor b_packed,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c_out,
int64_t M, int64_t K, int64_t L,
int64_t rest_m, int64_t rest_k
) {
auto options = torch::dtype(torch::kFloat16).device(a_packed.device());
auto b_decoded = torch::empty({L, K}, options);
decode_nvfp4_vector(
b_packed.data_ptr<uint8_t>(),
sfb.data_ptr<at::Half>(),
b_decoded.data_ptr<at::Half>(),
static_cast<int>(K),
static_cast<int>(L),
static_cast<int>(rest_k)
);
launch_nvfp4_kernel(
a_packed.data_ptr<uint8_t>(),
b_decoded.data_ptr<at::Half>(),
sfa.data_ptr<at::Half>(),
c_out.data_ptr<at::Half>(),
static_cast<int>(M),
static_cast<int>(K),
static_cast<int>(L),
static_cast<int>(rest_m),
static_cast<int>(rest_k)
);
return c_out;
}
"""
module = None
def get_module():
global module
if module is None:
module = load_inline(
name="nvfp4_gemv",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["nvfp4_gemv"],
verbose=False,
extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],
)
return module
def custom_kernel(data: input_t) -> output_t:
a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data
M, _, L = c_ref.shape
K = a_ref.shape[1] * 2
a_packed = a_ref.view(torch.uint8).contiguous()
b_packed = b_ref.view(torch.uint8).contiguous()
device = a_ref.device
sfa_scales = sfa_permuted.to(dtype=torch.float16, device=device, non_blocking=True).contiguous()
sfb_scales = sfb_permuted.to(dtype=torch.float16, device=device, non_blocking=True).contiguous()
rest_m = sfa_scales.shape[2]
rest_k = sfa_scales.shape[4]
if not c_ref.is_contiguous():
c_ref = c_ref.contiguous()
c_matrix = c_ref.view(M, L)
mod = get_module()
mod.nvfp4_gemv(a_packed, b_packed, sfa_scales, sfb_scales, c_matrix, M, K, L, rest_m, rest_k)
return c_ref
scrolls · 453 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