submission 109462
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 190 lines, June 9 Researcher Reciprocity License v1.0.
kmajor.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109462?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:6a89e5f3aa3863d80905f903f05219868991817ee1f797476877e1e31cdabace
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
vector-width = half2
__device__ __forceinline__ half2 process_16_elements(uint32_t a_lo, uint32_t a_hi, uint32_t b_lo, uint32_t b_hi, half2 scale_h2) {Kernel source
kmajor.py190 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 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__ 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)
);
}
__device__ __forceinline__ uint32_t cvt_f8x2_to_f16x2(uint16_t src) {
uint32_t dst;
asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(dst) : "h"(src));
return dst;
}
__device__ __forceinline__ half2 process_16_elements(uint32_t a_lo, uint32_t a_hi, uint32_t b_lo, uint32_t b_hi, half2 scale_h2) {
uint32_t a0, a1, a2, a3, a4, a5, a6, a7;
uint32_t b0, b1, b2, b3, b4, b5, b6, b7;
cvt_f4x8_to_f16x8(a_lo, a0, a1, a2, a3);
cvt_f4x8_to_f16x8(a_hi, a4, a5, a6, a7);
cvt_f4x8_to_f16x8(b_lo, b0, b1, b2, b3);
cvt_f4x8_to_f16x8(b_hi, b4, b5, b6, b7);
half2 ab0 = __hmul2(*reinterpret_cast<half2*>(&a0), *reinterpret_cast<half2*>(&b0));
half2 ab1 = __hmul2(*reinterpret_cast<half2*>(&a1), *reinterpret_cast<half2*>(&b1));
half2 ab2 = __hmul2(*reinterpret_cast<half2*>(&a2), *reinterpret_cast<half2*>(&b2));
half2 ab3 = __hmul2(*reinterpret_cast<half2*>(&a3), *reinterpret_cast<half2*>(&b3));
half2 ab4 = __hmul2(*reinterpret_cast<half2*>(&a4), *reinterpret_cast<half2*>(&b4));
half2 ab5 = __hmul2(*reinterpret_cast<half2*>(&a5), *reinterpret_cast<half2*>(&b5));
half2 ab6 = __hmul2(*reinterpret_cast<half2*>(&a6), *reinterpret_cast<half2*>(&b6));
half2 ab7 = __hmul2(*reinterpret_cast<half2*>(&a7), *reinterpret_cast<half2*>(&b7));
half2 sum01 = __hadd2(ab0, ab1);
half2 sum23 = __hadd2(ab2, ab3);
half2 sum45 = __hadd2(ab4, ab5);
half2 sum67 = __hadd2(ab6, ab7);
half2 sum0123 = __hadd2(sum01, sum23);
half2 sum4567 = __hadd2(sum45, sum67);
half2 local_sum = __hadd2(sum0123, sum4567);
return __hmul2(local_sum, scale_h2);
}
__global__ void gemv_kernel(
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_s2,
int64_t sfb_s2,
int64_t c_s0, int64_t c_s2
) {
constexpr int WARPS_PER_BLOCK = 4;
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;
// Scale factors are contiguous in K dimension (stride 1)
const uint8_t* sfa_row = SFA + row * sfa_s0 + batch * sfa_s2;
const uint8_t* sfb_batch = SFB + batch * sfb_s2;
const uint8_t* A_row = A + row * a_s0 + batch * a_s2;
const uint8_t* B_batch = B + batch * b_s2;
half2 acc_h2 = __float2half2_rn(0.0f);
int K_scales = K / 16;
for (int scale_base = 0; scale_base < K_scales; scale_base += 64) {
int lane_scale_base = scale_base + lane_id * 2;
if (lane_scale_base + 1 >= K_scales) {
// Tail handling
if (lane_scale_base < K_scales) {
uint8_t sfa0 = sfa_row[lane_scale_base];
uint8_t sfb0 = sfb_batch[lane_scale_base];
uint16_t sf_packed = (uint16_t(sfb0) << 8) | uint16_t(sfa0);
uint32_t sf_f16x2 = cvt_f8x2_to_f16x2(sf_packed);
half2 sf_h2 = *reinterpret_cast<half2*>(&sf_f16x2);
half scale = __hmul(sf_h2.x, sf_h2.y);
half2 scale_h2 = __half2half2(scale);
int k_byte = lane_scale_base * 8;
uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte);
uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte);
acc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.x, a_vec.y, b_vec.x, b_vec.y, scale_h2));
}
break;
}
// Vectorized scale factor loads - 2 bytes at once, coalesced across warp
uint16_t sfa_pair = *reinterpret_cast<const uint16_t*>(sfa_row + lane_scale_base);
uint16_t sfb_pair = *reinterpret_cast<const uint16_t*>(sfb_batch + lane_scale_base);
uint8_t sfa0 = sfa_pair & 0xFF;
uint8_t sfa1 = sfa_pair >> 8;
uint8_t sfb0 = sfb_pair & 0xFF;
uint8_t sfb1 = sfb_pair >> 8;
// Convert scales to fp16 and compute combined scale
uint16_t sf_packed0 = (uint16_t(sfb0) << 8) | uint16_t(sfa0);
uint16_t sf_packed1 = (uint16_t(sfb1) << 8) | uint16_t(sfa1);
uint32_t sf0_f16x2 = cvt_f8x2_to_f16x2(sf_packed0);
uint32_t sf1_f16x2 = cvt_f8x2_to_f16x2(sf_packed1);
half2 sf0_h2 = *reinterpret_cast<half2*>(&sf0_f16x2);
half2 sf1_h2 = *reinterpret_cast<half2*>(&sf1_f16x2);
half scale0 = __hmul(sf0_h2.x, sf0_h2.y);
half scale1 = __hmul(sf1_h2.x, sf1_h2.y);
half2 scale0_h2 = __half2half2(scale0);
half2 scale1_h2 = __half2half2(scale1);
// Load 16 bytes = 32 FP4 elements
int k_byte = lane_scale_base * 8;
uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_byte);
uint4 b_vec = *reinterpret_cast<const uint4*>(B_batch + k_byte);
// Process both scale groups
acc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.x, a_vec.y, b_vec.x, b_vec.y, scale0_h2));
acc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.z, a_vec.w, b_vec.z, b_vec.w, scale1_h2));
}
float acc = __half2float(acc_h2.x) + __half2float(acc_h2.y);
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);
constexpr int WARPS_PER_BLOCK = 4;
dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);
dim3 block(32 * WARPS_PER_BLOCK);
gemv_kernel<<<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(2),
sfb.stride(2),
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_vec_scales',
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, sfb, c)scrolls · 190 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 107336.
⋯ 59 unchanged linesreturn __hmul2(local_sum, scale_h2);}- __global__ void gemv_uint4(+ __global__ void gemv_kernel(const uint8_t* __restrict__ A,const uint8_t* __restrict__ B,const uint8_t* __restrict__ SFA,⋯ 1 unchanged lineshalf* __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 sfb_s1, int64_t sfb_s2,+ int64_t sfa_s0, int64_t sfa_s2,+ int64_t sfb_s2,int64_t c_s0, int64_t c_s2) {constexpr int WARPS_PER_BLOCK = 4;⋯ 4 unchanged linesif (row >= M) return;- int64_t sfa_base = row * sfa_s0 + batch * sfa_s2;- int64_t sfb_base = batch * sfb_s2;-+ // Scale factors are contiguous in K dimension (stride 1)+ const uint8_t* sfa_row = SFA + row * sfa_s0 + batch * sfa_s2;+ const uint8_t* sfb_batch = SFB + batch * sfb_s2;const uint8_t* A_row = A + row * a_s0 + batch * a_s2;const uint8_t* B_batch = B + batch * b_s2;half2 acc_h2 = __float2half2_rn(0.0f);int K_scales = K / 16;- // Each lane processes 2 scale groups (32 FP4 elements) per iteration- // 32 lanes * 2 groups = 64 scale groups per iterationfor (int scale_base = 0; scale_base < K_scales; scale_base += 64) {- int k_scale0 = scale_base + lane_id * 2;- int k_scale1 = k_scale0 + 1;+ int lane_scale_base = scale_base + lane_id * 2;- if (k_scale1 >= K_scales) {- // Handle tail: only process first group if valid- if (k_scale0 < K_scales) {- int sfa_idx = sfa_base + k_scale0 * sfa_s1;- int sfb_idx = sfb_base + k_scale0 * sfb_s1;- uint8_t sfa_val = SFA[sfa_idx];- uint8_t sfb_val = SFB[sfb_idx];- uint16_t sf_packed = (uint16_t(sfb_val) << 8) | uint16_t(sfa_val);+ if (lane_scale_base + 1 >= K_scales) {+ // Tail handling+ if (lane_scale_base < K_scales) {+ uint8_t sfa0 = sfa_row[lane_scale_base];+ uint8_t sfb0 = sfb_batch[lane_scale_base];+ uint16_t sf_packed = (uint16_t(sfb0) << 8) | uint16_t(sfa0);uint32_t sf_f16x2 = cvt_f8x2_to_f16x2(sf_packed);half2 sf_h2 = *reinterpret_cast<half2*>(&sf_f16x2);- half scale_h = __hmul(sf_h2.x, sf_h2.y);- half2 scale_h2 = __half2half2(scale_h);+ half scale = __hmul(sf_h2.x, sf_h2.y);+ half2 scale_h2 = __half2half2(scale);- int k_byte = k_scale0 * 8;+ int k_byte = lane_scale_base * 8;uint2 a_vec = *reinterpret_cast<const uint2*>(A_row + k_byte);uint2 b_vec = *reinterpret_cast<const uint2*>(B_batch + k_byte);acc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.x, a_vec.y, b_vec.x, b_vec.y, scale_h2));⋯ 1 unchanged linesbreak;}- // Load 16 bytes = 32 FP4 elements (2 scale groups)- int k_byte = k_scale0 * 8;- uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_byte);- uint4 b_vec = *reinterpret_cast<const uint4*>(B_batch + k_byte);+ // Vectorized scale factor loads - 2 bytes at once, coalesced across warp+ uint16_t sfa_pair = *reinterpret_cast<const uint16_t*>(sfa_row + lane_scale_base);+ uint16_t sfb_pair = *reinterpret_cast<const uint16_t*>(sfb_batch + lane_scale_base);- // Load 2 scale factors for A and B- int sfa_idx0 = sfa_base + k_scale0 * sfa_s1;- int sfa_idx1 = sfa_base + k_scale1 * sfa_s1;- int sfb_idx0 = sfb_base + k_scale0 * sfb_s1;- int sfb_idx1 = sfb_base + k_scale1 * sfb_s1;+ uint8_t sfa0 = sfa_pair & 0xFF;+ uint8_t sfa1 = sfa_pair >> 8;+ uint8_t sfb0 = sfb_pair & 0xFF;+ uint8_t sfb1 = sfb_pair >> 8;- uint8_t sfa0 = SFA[sfa_idx0], sfa1 = SFA[sfa_idx1];- uint8_t sfb0 = SFB[sfb_idx0], sfb1 = SFB[sfb_idx1];-- // Convert scales to fp16+ // Convert scales to fp16 and compute combined scaleuint16_t sf_packed0 = (uint16_t(sfb0) << 8) | uint16_t(sfa0);uint16_t sf_packed1 = (uint16_t(sfb1) << 8) | uint16_t(sfa1);uint32_t sf0_f16x2 = cvt_f8x2_to_f16x2(sf_packed0);⋯ 5 unchanged lineshalf2 scale0_h2 = __half2half2(scale0);half2 scale1_h2 = __half2half2(scale1);- // Process first 16 elements (first scale group)+ // Load 16 bytes = 32 FP4 elements+ int k_byte = lane_scale_base * 8;+ uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_byte);+ uint4 b_vec = *reinterpret_cast<const uint4*>(B_batch + k_byte);++ // Process both scale groupsacc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.x, a_vec.y, b_vec.x, b_vec.y, scale0_h2));- // Process second 16 elements (second scale group)acc_h2 = __hadd2(acc_h2, process_16_elements(a_vec.z, a_vec.w, b_vec.z, b_vec.w, scale1_h2));}⋯ 13 unchanged linesdim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L);dim3 block(32 * WARPS_PER_BLOCK);- gemv_uint4<<<grid, block>>>(+ gemv_kernel<<<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()),⋯ 1 unchanged linesreinterpret_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),- sfb.stride(1), sfb.stride(2),+ sfa.stride(0), sfa.stride(2),+ sfb.stride(2),c.stride(0), c.stride(2));return c;⋯ 3 unchanged linescpp_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_uint4',+ name='gemv_vec_scales',cpp_sources=cpp_source,cuda_sources=cuda_source,functions=['gemv_cuda'],
scrolls · 144 diff lines total
Best evidence level for this revision: reported
JSON