submission 112996
shigao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 411 lines, June 9 Researcher Reciprocity License v1.0.
node_12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-112996?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:51d9b146a2795ef5858034cc6a1c353239553098b6e9625ced4e6ed9dc2cc1c6
license declaredunknown
license concludedunknown
authorsshigao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
int vec_count = K / 32; // int4 count for packed FP4 (16B per chunk)fp8
__device__ __forceinline__ __nv_fp8_e4m3 byte_to_fp8(uint8_t v) {shared-memory
extern __shared__ uint8_t smem[];vector-width = int4
const int4* __restrict__ row_a_vec,Kernel source
node_12.py411 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import TypeVar
input_t = TypeVar("input_t", bound=tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor])
output_t = TypeVar("output_t", bound=torch.Tensor)
cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <ATen/cuda/Exceptions.h>
#include <ATen/cuda/CUDAContext.h>
#include <sm_61_intrinsics.h>
#include <vector>
// Precomputed FP4 (e2m1) → int8 lookup table for DP4A.
__device__ __align__(16) int PACKED_FP4_DP4A[65536];
constexpr float kFp4Rescale = 1.0f / 16.0f; // both operands scaled by 4
__device__ __forceinline__ int load_packed(uint16_t word) {
return __ldg(&PACKED_FP4_DP4A[word]);
}
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
return val;
}
__device__ __forceinline__ __nv_fp8_e4m3 byte_to_fp8(uint8_t v) {
union {
uint8_t u;
__nv_fp8_e4m3 fp;
} cvt;
cvt.u = v;
return cvt.fp;
}
__device__ __forceinline__ float accumulate_chunk_single(
int idx,
const int4* __restrict__ row_a_vec,
const int* __restrict__ row_b_packed,
const uchar2* __restrict__ sfa_pairs,
const float2* __restrict__ sfb_scales
) {
int4 va = __ldg(row_a_vec + idx);
const uint16_t* raw_a16 = reinterpret_cast<const uint16_t*>(&va);
uchar2 sfa_pair = __ldg(sfa_pairs + idx);
float2 sfb_scale = sfb_scales[idx];
float scale_0 = static_cast<float>(byte_to_fp8(sfa_pair.x)) * sfb_scale.x;
float scale_1 = static_cast<float>(byte_to_fp8(sfa_pair.y)) * sfb_scale.y;
int base = idx * 8;
int acc_int0 = 0;
#pragma unroll
for (int j = 0; j < 4; ++j) {
int b_val = row_b_packed[base + j];
acc_int0 = __dp4a(load_packed(raw_a16[j]), b_val, acc_int0);
}
int acc_int1 = 0;
#pragma unroll
for (int j = 4; j < 8; ++j) {
int b_val = row_b_packed[base + j];
acc_int1 = __dp4a(load_packed(raw_a16[j]), b_val, acc_int1);
}
float local0 = static_cast<float>(acc_int0) * kFp4Rescale;
float local1 = static_cast<float>(acc_int1) * kFp4Rescale;
return fmaf(local0, scale_0, local1 * scale_1);
}
__device__ __forceinline__ void accumulate_chunk_dual(
int idx,
const int4* __restrict__ row_a0_vec,
const int4* __restrict__ row_a1_vec,
const int* __restrict__ row_b_packed,
const uchar2* __restrict__ sfa0_pairs,
const uchar2* __restrict__ sfa1_pairs,
const float2* __restrict__ sfb_scales,
float& acc0,
float& acc1
) {
int4 va0 = __ldg(row_a0_vec + idx);
int4 va1 = __ldg(row_a1_vec + idx);
const uint16_t* raw_a16_0 = reinterpret_cast<const uint16_t*>(&va0);
const uint16_t* raw_a16_1 = reinterpret_cast<const uint16_t*>(&va1);
uchar2 sfa0_pair = __ldg(sfa0_pairs + idx);
uchar2 sfa1_pair = __ldg(sfa1_pairs + idx);
float2 sfb_scale = sfb_scales[idx];
float scale00 = static_cast<float>(byte_to_fp8(sfa0_pair.x)) * sfb_scale.x;
float scale01 = static_cast<float>(byte_to_fp8(sfa0_pair.y)) * sfb_scale.y;
float scale10 = static_cast<float>(byte_to_fp8(sfa1_pair.x)) * sfb_scale.x;
float scale11 = static_cast<float>(byte_to_fp8(sfa1_pair.y)) * sfb_scale.y;
int base = idx * 8;
int acc0_0 = 0, acc0_1 = 0;
int acc1_0 = 0, acc1_1 = 0;
#pragma unroll
for (int j = 0; j < 4; ++j) {
int b_val = row_b_packed[base + j];
acc0_0 = __dp4a(load_packed(raw_a16_0[j]), b_val, acc0_0);
acc1_0 = __dp4a(load_packed(raw_a16_1[j]), b_val, acc1_0);
}
#pragma unroll
for (int j = 4; j < 8; ++j) {
int b_val = row_b_packed[base + j];
acc0_1 = __dp4a(load_packed(raw_a16_0[j]), b_val, acc0_1);
acc1_1 = __dp4a(load_packed(raw_a16_1[j]), b_val, acc1_1);
}
float scaled0 = static_cast<float>(acc0_0) * kFp4Rescale;
float scaled1 = static_cast<float>(acc0_1) * kFp4Rescale;
acc0 = fmaf(scaled0, scale00, scaled1 * scale01) + acc0;
float scaled2 = static_cast<float>(acc1_0) * kFp4Rescale;
float scaled3 = static_cast<float>(acc1_1) * kFp4Rescale;
acc1 = fmaf(scaled2, scale10, scaled3 * scale11) + acc1;
}
template<int WARPS_PER_BLOCK>
__global__ void nvfp4_gemv_shared_kernel(
const uint8_t* __restrict__ A, // [M, K/2, L]
const uint8_t* __restrict__ B, // [1, K/2, L] (only first row used)
const __nv_fp8_e4m3* __restrict__ SFA, // [M, K/16, L]
const __nv_fp8_e4m3* __restrict__ SFB, // [1, K/16, L]
half* __restrict__ C, // [M, 1, L]
int M, int K, int L,
int64_t stride_a_m, int64_t stride_a_l,
int64_t stride_b_l,
int64_t stride_sfa_m, int64_t stride_sfa_l,
int64_t stride_sfb_l,
int64_t stride_c_m, int64_t stride_c_l
) {
int tid = threadIdx.x;
int warp_id = tid >> 5;
int lane = tid & 31;
int base_m = blockIdx.x * (WARPS_PER_BLOCK * 2);
int l = blockIdx.y;
if (base_m >= M || l >= L) {
return;
}
extern __shared__ uint8_t smem[];
int packed_words = K >> 2; // K/4 packed int32 entries
int b_packed_bytes = packed_words * sizeof(int);
int b_packed_aligned = (b_packed_bytes + 15) & ~15;
uint8_t* smem_b = smem;
int* shared_b_packed = reinterpret_cast<int*>(smem_b);
float2* shared_sfb_scale = reinterpret_cast<float2*>(smem_b + b_packed_aligned);
// Cache packed B and SFB for this batch l once per block.
const uint8_t* b_global = B + l * stride_b_l;
const __nv_fp8_e4m3* sfb_global = SFB + l * stride_sfb_l;
const uint16_t* b_words_global = reinterpret_cast<const uint16_t*>(b_global);
for (int idx = tid; idx < packed_words; idx += blockDim.x) {
uint16_t w = __ldg(b_words_global + idx);
shared_b_packed[idx] = load_packed(w);
}
int sfb_pairs = K / 32; // two fp8 per pair
const uchar2* sfb_vec_global = reinterpret_cast<const uchar2*>(sfb_global);
for (int idx = tid; idx < sfb_pairs; idx += blockDim.x) {
uchar2 v = sfb_vec_global[idx];
float2 scale_pair;
scale_pair.x = static_cast<float>(byte_to_fp8(v.x));
scale_pair.y = static_cast<float>(byte_to_fp8(v.y));
shared_sfb_scale[idx] = scale_pair;
}
__syncthreads();
if (warp_id >= WARPS_PER_BLOCK) return;
int m0 = base_m + warp_id * 2;
int m1 = m0 + 1;
bool has_row0 = m0 < M;
bool has_row1 = m1 < M;
if (!has_row0) return;
const uint8_t* row_a0 = A + m0 * stride_a_m + l * stride_a_l;
const __nv_fp8_e4m3* row_sfa0 = SFA + m0 * stride_sfa_m + l * stride_sfa_l;
const int4* row_a0_vec = reinterpret_cast<const int4*>(row_a0);
const uchar2* sfa0_pairs = reinterpret_cast<const uchar2*>(row_sfa0);
const int* row_b_packed = shared_b_packed;
const float2* sfb_scales = shared_sfb_scale;
int vec_count = K / 32; // int4 count for packed FP4 (16B per chunk)
if (has_row1) {
const uint8_t* row_a1 = A + m1 * stride_a_m + l * stride_a_l;
const __nv_fp8_e4m3* row_sfa1 = SFA + m1 * stride_sfa_m + l * stride_sfa_l;
const int4* row_a1_vec = reinterpret_cast<const int4*>(row_a1);
const uchar2* sfa1_pairs = reinterpret_cast<const uchar2*>(row_sfa1);
float acc_row0 = 0.0f;
float acc_row1 = 0.0f;
// Double-stripmine to reduce loop overhead while keeping ILP.
for (int idx = lane; idx < vec_count; idx += 64) {
accumulate_chunk_dual(idx, row_a0_vec, row_a1_vec, row_b_packed, sfa0_pairs, sfa1_pairs, sfb_scales, acc_row0, acc_row1);
int idx_next = idx + 32;
if (idx_next < vec_count) {
accumulate_chunk_dual(idx_next, row_a0_vec, row_a1_vec, row_b_packed, sfa0_pairs, sfa1_pairs, sfb_scales, acc_row0, acc_row1);
}
}
acc_row0 = warp_reduce_sum(acc_row0);
acc_row1 = warp_reduce_sum(acc_row1);
if (lane == 0) {
int64_t out0 = static_cast<int64_t>(m0) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
C[out0] = __float2half(acc_row0);
int64_t out1 = static_cast<int64_t>(m1) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
C[out1] = __float2half(acc_row1);
}
} else {
float acc = 0.0f;
for (int idx = lane; idx < vec_count; idx += 64) {
acc += accumulate_chunk_single(idx, row_a0_vec, row_b_packed, sfa0_pairs, sfb_scales);
int idx_next = idx + 32;
if (idx_next < vec_count) {
acc += accumulate_chunk_single(idx_next, row_a0_vec, row_b_packed, sfa0_pairs, sfb_scales);
}
}
acc = warp_reduce_sum(acc);
if (lane == 0) {
int64_t out_idx = static_cast<int64_t>(m0) * stride_c_m + static_cast<int64_t>(l) * stride_c_l;
C[out_idx] = __float2half(acc);
}
}
}
template<int WARPS_PER_BLOCK>
void launch_gemv(
const uint8_t* A, const uint8_t* B,
const __nv_fp8_e4m3* SFA, const __nv_fp8_e4m3* SFB,
half* C,
int M, int K, int L,
int64_t stride_a_m, int64_t stride_a_l,
int64_t stride_b_l,
int64_t stride_sfa_m, int64_t stride_sfa_l,
int64_t stride_sfb_l,
int64_t stride_c_m, int64_t stride_c_l,
int shared_bytes
) {
dim3 block(32 * WARPS_PER_BLOCK);
dim3 grid((M + (2 * WARPS_PER_BLOCK) - 1) / (2 * WARPS_PER_BLOCK), L);
nvfp4_gemv_shared_kernel<WARPS_PER_BLOCK><<<grid, block, shared_bytes>>>(
A, B, SFA, SFB, C,
M, K, L,
stride_a_m, stride_a_l,
stride_b_l,
stride_sfa_m, stride_sfa_l,
stride_sfb_l,
stride_c_m, stride_c_l
);
}
void nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c
) {
int M = a.size(0);
int K = a.size(1) * 2; // logical elements
int L = a.size(2);
TORCH_CHECK(K % 64 == 0, "K must be divisible by 64");
static bool lut_initialized = false;
if (!lut_initialized) {
static const int8_t E2M1_SCALED[16] = {
0, 2, 4, 6, 8, 12, 16, 24,
0, -2, -4, -6, -8, -12, -16, -24
};
std::vector<int> pack_lut(65536);
for (int word = 0; word < 65536; ++word) {
uint16_t v = static_cast<uint16_t>(word);
int8_t e0 = E2M1_SCALED[v & 0x0F];
int8_t e1 = E2M1_SCALED[(v >> 4) & 0x0F];
int8_t e2 = E2M1_SCALED[(v >> 8) & 0x0F];
int8_t e3 = E2M1_SCALED[(v >> 12) & 0x0F];
uint32_t packed =
static_cast<uint8_t>(e0) |
(static_cast<uint8_t>(e1) << 8) |
(static_cast<uint8_t>(e2) << 16) |
(static_cast<uint8_t>(e3) << 24);
pack_lut[word] = static_cast<int>(packed);
}
cudaMemcpyToSymbol(PACKED_FP4_DP4A, pack_lut.data(), pack_lut.size() * sizeof(int));
lut_initialized = true;
}
int b_pack_bytes = K; // K/4 int entries → K bytes
int sfb_scale_bytes = (K >> 5) * static_cast<int>(sizeof(float2)); // K/32 pairs of float2 scales
int shared_bytes = ((b_pack_bytes + 15) & ~15) + ((sfb_scale_bytes + 15) & ~15);
const uint8_t* a_ptr = (const uint8_t*)a.data_ptr();
const uint8_t* b_ptr = (const uint8_t*)b.data_ptr();
const __nv_fp8_e4m3* sfa_ptr = (const __nv_fp8_e4m3*)sfa.data_ptr();
const __nv_fp8_e4m3* sfb_ptr = (const __nv_fp8_e4m3*)sfb.data_ptr();
half* c_ptr = (half*)c.data_ptr();
// Warp selection tuned for dual-row processing.
// Use 8 warps for medium/large K to keep occupancy healthy; 4 for very small.
if (K <= 4096) {
launch_gemv<4>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_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),
shared_bytes
);
} else {
launch_gemv<8>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_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),
shared_bytes
);
}
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
printf("Kernel failed: %s\n", cudaGetErrorString(err));
}
}
"""
cpp_source = """
void nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c
);
"""
nvfp4_gemv_module = load_inline(
name='nvfp4_gemv',
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
functions=['nvfp4_gemv'],
verbose=True,
extra_cflags=[
'-std=c++17',
"-O3", # CPU 代码最高优化
"-march=native", # 🚀 针对本机 CPU 指令集优化 (减少 launch overhead)
"-fno-math-errno", # 移除数学错误检查
"-Wall", # 打开警告,防止低级错误
],
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
'--extra-device-vectorization',
"--restrict",
"-std=c++17",
"--ptxas-options=-O3",
"--expt-relaxed-constexpr",
"-arch=sm_100a",
'-Xptxas', '-v',
'-lineinfo',
'-U__CUDA_NO_HALF_OPERATORS__', # Enable half operators
'-U__CUDA_NO_HALF_CONVERSIONS__', # Enable half conversions
],
)
def custom_kernel(data: input_t) -> output_t:
a_ref, b_ref, sfa, sfb, _, _, c_ref = data
nvfp4_gemv_module.nvfp4_gemv(
a_ref,
b_ref,
sfa,
sfb,
c_ref
)
return c_refscrolls · 411 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