submission 116208
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 457 lines, June 9 Researcher Reciprocity License v1.0.
template.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-116208?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:458c46695ac7c4007f3b26674b8125b2b1608b8ad377f579b26022ebb7dd682c
license declaredunknown
license concludedunknown
authorsshiyegao
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
template.py457 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 <ATen/core/dispatch/Dispatcher.h>
#include <sm_61_intrinsics.h>
#include <vector>
// Helper: match the blocking layout used by torch._scaled_mm reference path.
static inline torch::Tensor to_blocked(const torch::Tensor& input) {
// input shape: [rows, cols], rows divisible by 128, cols divisible by 4.
auto rows = input.size(0);
auto cols = input.size(1);
auto n_row_blocks = (rows + 127) / 128;
auto n_col_blocks = (cols + 3) / 4;
auto padded = input.contiguous();
auto blocks = padded.view({n_row_blocks, 128, n_col_blocks, 4}).permute({0, 2, 1, 3});
auto rearranged = blocks.reshape({-1, 4, 32, 4}).transpose(1, 2).reshape({-1, 32, 16});
return rearranged.reshape({-1});
}
// Core idea:
// 1) Fast path: use torch::_scaled_mm with blocked scale factors and N padded to 128,
// mapping to the FP4 tensor core kernel on Blackwell/RTX 5090. This mirrors the
// reference layout to hit the optimized path and keep correctness.
// 2) Fallback: DP4A implementation kept from previous nodes to preserve compatibility
// if _scaled_mm is unavailable at runtime.
// ==== DP4A fallback pieces ====
__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 uchar2* __restrict__ sfb_pairs
) {
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);
uchar2 sfb_pair = __ldg(sfb_pairs + idx);
float scale_0 = static_cast<float>(byte_to_fp8(sfa_pair.x)) * static_cast<float>(byte_to_fp8(sfb_pair.x));
float scale_1 = static_cast<float>(byte_to_fp8(sfa_pair.y)) * static_cast<float>(byte_to_fp8(sfb_pair.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 uchar2* __restrict__ sfb_pairs,
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);
uchar2 sfb_pair = __ldg(sfb_pairs + idx);
float scale00 = static_cast<float>(byte_to_fp8(sfa0_pair.x)) * static_cast<float>(byte_to_fp8(sfb_pair.x));
float scale01 = static_cast<float>(byte_to_fp8(sfa0_pair.y)) * static_cast<float>(byte_to_fp8(sfb_pair.y));
float scale10 = static_cast<float>(byte_to_fp8(sfa1_pair.x)) * static_cast<float>(byte_to_fp8(sfb_pair.x));
float scale11 = static_cast<float>(byte_to_fp8(sfa1_pair.y)) * static_cast<float>(byte_to_fp8(sfb_pair.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, // [128, K/2, L] (only first row needed)
const __nv_fp8_e4m3* __restrict__ SFA, // [M, K/16, L]
const __nv_fp8_e4m3* __restrict__ SFB, // [128, 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
uint8_t* smem_b = smem;
int* shared_b_packed = reinterpret_cast<int*>(smem_b);
// Cache packed B for this batch l once per block; SFB stays in global to trim smem.
const uint8_t* b_global = B + l * stride_b_l;
const __nv_fp8_e4m3* sfb_global = SFB + l * stride_sfb_l;
const uchar2* sfb_pairs_global = reinterpret_cast<const uchar2*>(sfb_global);
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);
}
__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 uchar2* sfb_pairs = sfb_pairs_global;
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_pairs, 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_pairs, 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_pairs);
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_pairs);
}
}
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
);
}
// ==== Front-end dispatch ====
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");
// ----- Fast path: tensor-core _scaled_mm with blocked scales (N padded to 128) -----
static bool scaled_mm_supported = true;
if (scaled_mm_supported) {
static auto scaled_mm_call = c10::Dispatcher::singleton()
.findSchemaOrThrow("aten::_scaled_mm", "")
.typed<torch::Tensor(
const torch::Tensor&,
const torch::Tensor&,
const torch::Tensor&,
const torch::Tensor&,
const c10::optional<torch::Tensor>&,
const c10::optional<torch::Tensor>&,
c10::optional<torch::ScalarType>,
bool
)>();
try {
for (int l = 0; l < L; ++l) {
auto a_slice = a.select(2, l).contiguous(); // [M, K/2]
auto b_mat = b.select(2, l).transpose(0, 1).contiguous(); // [K/2, 128]
auto scale_a = to_blocked(sfa.select(2, l));
auto scale_b = to_blocked(sfb.select(2, l));
auto out = scaled_mm_call.call(
a_slice,
b_mat,
scale_a,
scale_b,
c10::optional<torch::Tensor>(),
c10::optional<torch::Tensor>(),
c10::make_optional(torch::kFloat16),
true // fast accumulation to favor perf on tensor cores
);
c.select(2, l).copy_(out.select(1, 0));
}
return;
} catch (const c10::Error&) {
scaled_mm_supported = false;
}
}
// ----- Fallback: DP4A kernel -----
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 shared_bytes = ((b_pack_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();
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 · 457 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