submission 70215
lyi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 553 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70215?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:89b2a69fd381fdefab22cfec784eb2eb232c0831250b435d84f3a3c9e0a32f2e
license declaredunknown
license concludedunknown
authorslyi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float shared_scale_b[K_BLOCK_TILE];Kernel source
submission.py553 lines
import torch
from torch.utils.cpp_extension import load_inline
import os
from task import input_t, output_t
_MODULE = None
_BYTES_PER_BLOCK = 8
_PROFILE_ENV = "NVFP4_PROFILE"
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <math.h>
__device__ __constant__ float kFp4Lut[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_e4m3(uint8_t value) {
const int sign = (value & 0x80) ? -1 : 1;
const int exp = (value >> 3) & 0x0F;
const int mant = value & 0x7;
if (exp == 0) {
return 0.0f;
}
if (exp == 0x0F) {
return sign * INFINITY;
}
const float base = 1.0f + static_cast<float>(mant) / 8.0f;
const int exp_unbiased = exp - 7;
return sign * ldexpf(base, exp_unbiased);
}
__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;
}
constexpr int kValuesPerBlock = 16;
constexpr int kBytesPerBlock = 8;
template <int ROWS_PER_CTA, int K_BLOCK_TILE>
__global__ void nvfp4_gemv_kernel_rows(
const uint8_t* __restrict__ a,
const uint8_t* __restrict__ b,
const uint8_t* __restrict__ scale_a_perm,
const uint8_t* __restrict__ scale_b_perm,
__half* __restrict__ out,
const int64_t m,
const int64_t k_blocks,
const int64_t stride_a0,
const int64_t stride_a1,
const int64_t stride_a2,
const int64_t stride_b0,
const int64_t stride_b1,
const int64_t stride_out0,
const int64_t stride_out1,
const int n32,
const int n4,
const int nblk,
const int64_t sfa_stride0,
const int64_t sfa_stride1,
const int64_t sfa_stride2,
const int64_t sfa_stride3,
const int64_t sfa_stride4,
const int64_t sfa_stride5,
const int64_t sfb_stride0,
const int64_t sfb_stride1,
const int64_t sfb_stride2,
const int64_t sfb_stride3,
const int64_t sfb_stride4,
const int64_t sfb_stride5
) {
__shared__ float shared_scale_b[K_BLOCK_TILE];
__shared__ float shared_b_vals[K_BLOCK_TILE * kValuesPerBlock];
const int row_lane = threadIdx.x;
const int64_t row = blockIdx.x * ROWS_PER_CTA + row_lane;
const int64_t batch = blockIdx.y;
const bool row_active = row < m;
const uint8_t* a_row = row_active ? (a + row * stride_a0 + batch * stride_a2) : nullptr;
__half* out_ptr = row_active ? (out + row * stride_out0 + batch * stride_out1) : nullptr;
const uint8_t* b_batch = b + batch * stride_b1;
const uint8_t* sfa_batch = scale_a_perm + batch * sfa_stride5;
const uint8_t* sfb_batch = scale_b_perm + batch * sfb_stride5;
int mm32 = 0, mm4 = 0, mm = 0;
if (row_active) {
mm32 = static_cast<int>(row & 31);
mm4 = static_cast<int>((row >> 5) & 3);
mm = static_cast<int>(row >> 7);
}
const uint8_t* sfa_row_base = row_active
? (sfa_batch + mm32 * sfa_stride0 + mm4 * sfa_stride1 + mm * sfa_stride2)
: nullptr;
const uint8_t* sfb_base = sfb_batch + n32 * sfb_stride0 + n4 * sfb_stride1 + nblk * sfb_stride2;
float acc = 0.0f;
for (int64_t tile_block = 0; tile_block < k_blocks; tile_block += K_BLOCK_TILE) {
const int64_t remaining = k_blocks - tile_block;
const int blocks_here = remaining > K_BLOCK_TILE ? K_BLOCK_TILE : static_cast<int>(remaining);
const int tile_values = blocks_here * kValuesPerBlock;
for (int idx = row_lane; idx < blocks_here; idx += ROWS_PER_CTA) {
const int global_block = static_cast<int>(tile_block) + idx;
const int kk4 = global_block & 3;
const int kk = global_block >> 2;
const uint8_t* scale_ptr = sfb_base + kk4 * sfb_stride3 + kk * sfb_stride4;
shared_scale_b[idx] = decode_fp8_e4m3(*scale_ptr);
}
__syncthreads();
for (int idx = row_lane; idx < tile_values; idx += ROWS_PER_CTA) {
const int blk_local = idx / kValuesPerBlock;
const int nib_idx = idx - blk_local * kValuesPerBlock;
const int byte_in_block = nib_idx >> 1;
const bool hi = nib_idx & 1;
const int64_t global_block = tile_block + blk_local;
const int64_t byte_index = global_block * kBytesPerBlock + byte_in_block;
const uint8_t byte_val = b_batch[byte_index * stride_b0];
const uint8_t nib = hi ? (byte_val >> 4) : (byte_val & 0xF);
const float scaled = kFp4Lut[nib] * shared_scale_b[blk_local];
shared_b_vals[blk_local * kValuesPerBlock + nib_idx] = scaled;
}
__syncthreads();
if (row_active) {
for (int blk_local = 0; blk_local < blocks_here; ++blk_local) {
const int64_t global_block = tile_block + blk_local;
const int kk4 = static_cast<int>(global_block & 3);
const int kk = static_cast<int>(global_block >> 2);
const uint8_t* scale_ptr = sfa_row_base + kk4 * sfa_stride3 + kk * sfa_stride4;
const float scale_a_val = decode_fp8_e4m3(*scale_ptr);
const int64_t block_byte_base = global_block * kBytesPerBlock;
#pragma unroll
for (int byte_offset = 0; byte_offset < kBytesPerBlock; ++byte_offset) {
const int64_t byte_index = block_byte_base + byte_offset;
const uint8_t a_byte = a_row[byte_index * stride_a1];
const float aval_lo = kFp4Lut[a_byte & 0xF] * scale_a_val;
const float aval_hi = kFp4Lut[(a_byte >> 4) & 0xF] * scale_a_val;
const int nib_base = blk_local * kValuesPerBlock + byte_offset * 2;
const float b_lo = shared_b_vals[nib_base];
const float b_hi = shared_b_vals[nib_base + 1];
acc = fmaf(aval_lo, b_lo, acc);
acc = fmaf(aval_hi, b_hi, acc);
}
}
}
__syncthreads();
}
if (row_active) {
out_ptr[0] = __float2half(acc);
}
}
template <int WARPS_PER_CTA, int K_BLOCK_TILE>
__global__ void nvfp4_gemv_kernel_warp(
const uint8_t* __restrict__ a,
const uint8_t* __restrict__ b,
const uint8_t* __restrict__ scale_a_perm,
const uint8_t* __restrict__ scale_b_perm,
__half* __restrict__ out,
const int64_t m,
const int64_t k_blocks,
const int64_t stride_a0,
const int64_t stride_a1,
const int64_t stride_a2,
const int64_t stride_b0,
const int64_t stride_b1,
const int64_t stride_out0,
const int64_t stride_out1,
const int n32,
const int n4,
const int nblk,
const int64_t sfa_stride0,
const int64_t sfa_stride1,
const int64_t sfa_stride2,
const int64_t sfa_stride3,
const int64_t sfa_stride4,
const int64_t sfa_stride5,
const int64_t sfb_stride0,
const int64_t sfb_stride1,
const int64_t sfb_stride2,
const int64_t sfb_stride3,
const int64_t sfb_stride4,
const int64_t sfb_stride5
) {
__shared__ float shared_scale_b[K_BLOCK_TILE];
__shared__ float shared_b_vals[K_BLOCK_TILE * kValuesPerBlock];
constexpr int THREADS_PER_CTA = WARPS_PER_CTA * 32;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int64_t row = blockIdx.x * WARPS_PER_CTA + warp;
const int64_t batch = blockIdx.y;
const bool row_active = row < m;
const uint8_t* a_row = row_active ? (a + row * stride_a0 + batch * stride_a2) : nullptr;
__half* out_ptr = row_active ? (out + row * stride_out0 + batch * stride_out1) : nullptr;
const uint8_t* b_batch = b + batch * stride_b1;
const uint8_t* sfa_batch = scale_a_perm + batch * sfa_stride5;
const uint8_t* sfb_batch = scale_b_perm + batch * sfb_stride5;
int mm32 = 0, mm4 = 0, mm = 0;
if (row_active) {
mm32 = static_cast<int>(row & 31);
mm4 = static_cast<int>((row >> 5) & 3);
mm = static_cast<int>(row >> 7);
}
const uint8_t* sfa_row_base = row_active
? (sfa_batch + mm32 * sfa_stride0 + mm4 * sfa_stride1 + mm * sfa_stride2)
: nullptr;
const uint8_t* sfb_base = sfb_batch + n32 * sfb_stride0 + n4 * sfb_stride1 + nblk * sfb_stride2;
float acc_lane = 0.0f;
for (int64_t tile_block = 0; tile_block < k_blocks; tile_block += K_BLOCK_TILE) {
const int64_t remaining = k_blocks - tile_block;
const int blocks_here = remaining > K_BLOCK_TILE ? K_BLOCK_TILE : static_cast<int>(remaining);
const int tile_values = blocks_here * kValuesPerBlock;
for (int idx = threadIdx.x; idx < blocks_here; idx += THREADS_PER_CTA) {
const int global_block = static_cast<int>(tile_block) + idx;
const int kk4 = global_block & 3;
const int kk = global_block >> 2;
const uint8_t* scale_ptr = sfb_base + kk4 * sfb_stride3 + kk * sfb_stride4;
shared_scale_b[idx] = decode_fp8_e4m3(*scale_ptr);
}
__syncthreads();
for (int idx = threadIdx.x; idx < tile_values; idx += THREADS_PER_CTA) {
const int blk_local = idx / kValuesPerBlock;
const int nib_idx = idx - blk_local * kValuesPerBlock;
const int byte_in_block = nib_idx >> 1;
const bool hi = nib_idx & 1;
const int64_t global_block = tile_block + blk_local;
const int64_t byte_index = global_block * kBytesPerBlock + byte_in_block;
const uint8_t byte_val = b_batch[byte_index * stride_b0];
const uint8_t nib = hi ? (byte_val >> 4) : (byte_val & 0xF);
const float scaled = kFp4Lut[nib] * shared_scale_b[blk_local];
shared_b_vals[blk_local * kValuesPerBlock + nib_idx] = scaled;
}
__syncthreads();
if (row_active) {
for (int blk_local = 0; blk_local < blocks_here; ++blk_local) {
const int64_t global_block = tile_block + blk_local;
const int kk4 = static_cast<int>(global_block & 3);
const int kk = static_cast<int>(global_block >> 2);
const uint8_t* scale_ptr = sfa_row_base + kk4 * sfa_stride3 + kk * sfa_stride4;
const float scale_a_val = decode_fp8_e4m3(*scale_ptr);
const int64_t block_byte_base = global_block * kBytesPerBlock;
for (int byte_offset = lane; byte_offset < kBytesPerBlock; byte_offset += 32) {
const int64_t byte_index = block_byte_base + byte_offset;
const uint8_t a_byte = a_row[byte_index * stride_a1];
const float aval_lo = kFp4Lut[a_byte & 0xF] * scale_a_val;
const float aval_hi = kFp4Lut[(a_byte >> 4) & 0xF] * scale_a_val;
const int nib_base = blk_local * kValuesPerBlock + byte_offset * 2;
const float b_lo = shared_b_vals[nib_base];
const float b_hi = shared_b_vals[nib_base + 1];
acc_lane = fmaf(aval_lo, b_lo, acc_lane);
acc_lane = fmaf(aval_hi, b_hi, acc_lane);
}
}
}
__syncthreads();
}
if (row_active) {
float acc = warp_reduce_sum(acc_lane);
if (lane == 0) {
out_ptr[0] = __float2half(acc);
}
}
}
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor scale_a_perm,
torch::Tensor scale_b_perm,
torch::Tensor out
) {
TORCH_CHECK(a.device().is_cuda(), "tensor a must be on CUDA");
TORCH_CHECK(b.device().is_cuda(), "tensor b must be on CUDA");
TORCH_CHECK(scale_a_perm.device().is_cuda(), "scale_a_perm must be on CUDA");
TORCH_CHECK(scale_b_perm.device().is_cuda(), "scale_b_perm must be on CUDA");
TORCH_CHECK(out.device().is_cuda(), "output tensor must be on CUDA");
TORCH_CHECK(a.scalar_type() == at::kByte, "a must be uint8 view");
TORCH_CHECK(b.scalar_type() == at::kByte, "b must be uint8 view");
TORCH_CHECK(scale_a_perm.scalar_type() == at::kByte, "scale_a_perm must be uint8 view");
TORCH_CHECK(scale_b_perm.scalar_type() == at::kByte, "scale_b_perm must be uint8 view");
TORCH_CHECK(out.scalar_type() == at::kHalf, "out must be float16");
TORCH_CHECK(a.dim() == 3, "a must be [M, K/2, L]");
TORCH_CHECK(b.dim() == 2, "b must be [K/2, L]");
TORCH_CHECK(scale_a_perm.dim() == 6, "scale_a_perm must be 6D");
TORCH_CHECK(scale_b_perm.dim() == 6, "scale_b_perm must be 6D");
TORCH_CHECK(out.dim() == 2, "out must be [M, L]");
const int64_t m = a.size(0);
const int64_t k_packed = a.size(1);
const int64_t l = a.size(2);
TORCH_CHECK(k_packed % kBytesPerBlock == 0, "K dimension must align to 16 elements");
const int64_t k_blocks = k_packed / kBytesPerBlock;
TORCH_CHECK(b.size(0) == k_packed && b.size(1) == l, "b shape mismatch");
TORCH_CHECK(scale_a_perm.size(5) == l, "scale_a_perm batch mismatch");
TORCH_CHECK(scale_b_perm.size(5) == l, "scale_b_perm batch mismatch");
TORCH_CHECK(out.size(0) == m && out.size(1) == l, "out shape mismatch");
at::cuda::CUDAGuard guard(a.device());
const auto stride_a0 = a.stride(0);
const auto stride_a1 = a.stride(1);
const auto stride_a2 = a.stride(2);
const auto stride_b0 = b.stride(0);
const auto stride_b1 = b.stride(1);
const auto stride_out0 = out.stride(0);
const auto stride_out1 = out.stride(1);
const auto sfa_stride0 = scale_a_perm.stride(0);
const auto sfa_stride1 = scale_a_perm.stride(1);
const auto sfa_stride2 = scale_a_perm.stride(2);
const auto sfa_stride3 = scale_a_perm.stride(3);
const auto sfa_stride4 = scale_a_perm.stride(4);
const auto sfa_stride5 = scale_a_perm.stride(5);
const auto sfb_stride0 = scale_b_perm.stride(0);
const auto sfb_stride1 = scale_b_perm.stride(1);
const auto sfb_stride2 = scale_b_perm.stride(2);
const auto sfb_stride3 = scale_b_perm.stride(3);
const auto sfb_stride4 = scale_b_perm.stride(4);
const auto sfb_stride5 = scale_b_perm.stride(5);
constexpr int kRowsPerCta = 128;
constexpr int kRowBlockTile = 128;
constexpr int kWarpCtas = 4;
constexpr int kWarpBlockTile = 64;
const int n_index = 0;
const int n32 = n_index & 31;
const int n4 = (n_index >> 5) & 3;
const int nblk = n_index >> 7;
auto stream = at::cuda::getCurrentCUDAStream();
const bool use_warp_kernel = (l <= 2) && (k_blocks >= 512);
if (use_warp_kernel) {
const dim3 block_dim(kWarpCtas * 32);
const dim3 grid_dim((m + kWarpCtas - 1) / kWarpCtas, l);
nvfp4_gemv_kernel_warp<kWarpCtas, kWarpBlockTile><<<grid_dim, block_dim, 0, stream>>>(
a.data_ptr<uint8_t>(),
b.data_ptr<uint8_t>(),
scale_a_perm.data_ptr<uint8_t>(),
scale_b_perm.data_ptr<uint8_t>(),
reinterpret_cast<__half*>(out.data_ptr<at::Half>()),
m,
k_blocks,
stride_a0,
stride_a1,
stride_a2,
stride_b0,
stride_b1,
stride_out0,
stride_out1,
n32,
n4,
nblk,
sfa_stride0,
sfa_stride1,
sfa_stride2,
sfa_stride3,
sfa_stride4,
sfa_stride5,
sfb_stride0,
sfb_stride1,
sfb_stride2,
sfb_stride3,
sfb_stride4,
sfb_stride5
);
} else {
const dim3 block_dim(kRowsPerCta);
const dim3 grid_dim((m + kRowsPerCta - 1) / kRowsPerCta, l);
nvfp4_gemv_kernel_rows<kRowsPerCta, kRowBlockTile><<<grid_dim, block_dim, 0, stream>>>(
a.data_ptr<uint8_t>(),
b.data_ptr<uint8_t>(),
scale_a_perm.data_ptr<uint8_t>(),
scale_b_perm.data_ptr<uint8_t>(),
reinterpret_cast<__half*>(out.data_ptr<at::Half>()),
m,
k_blocks,
stride_a0,
stride_a1,
stride_a2,
stride_b0,
stride_b1,
stride_out0,
stride_out1,
n32,
n4,
nblk,
sfa_stride0,
sfa_stride1,
sfa_stride2,
sfa_stride3,
sfa_stride4,
sfa_stride5,
sfb_stride0,
sfb_stride1,
sfb_stride2,
sfb_stride3,
sfb_stride4,
sfb_stride5
);
}
auto cuda_status = cudaGetLastError();
TORCH_CHECK(cuda_status == cudaSuccess, "nvfp4_gemv kernel launch failed: ", cudaGetErrorString(cuda_status));
return out;
}
"""
_CPP_SRC = r"""
#include <torch/extension.h>
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor scale_a_perm,
torch::Tensor scale_b_perm,
torch::Tensor out
);
"""
def _load_module():
global _MODULE
if _MODULE is None:
_MODULE = load_inline(
name="nvfp4_gemv_inline_kernel",
cpp_sources=_CPP_SRC,
cuda_sources=_CUDA_SRC,
functions=["nvfp4_gemv"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
return _MODULE
def _prepare_matrix_bytes(tensor: torch.Tensor) -> torch.Tensor:
return tensor.view(torch.uint8)
def _prepare_vector_bytes(tensor: torch.Tensor) -> torch.Tensor:
if tensor.size(0) == 0:
raise ValueError("Input vector has empty N dimension")
bytes_view = tensor.view(torch.uint8)
vec = bytes_view.select(0, 0)
return vec if vec.is_contiguous() else vec.contiguous()
def _prepare_perm_bytes(tensor: torch.Tensor, device: torch.device) -> torch.Tensor:
if tensor.device != device:
tensor = tensor.to(device=device, non_blocking=True)
if not tensor.is_contiguous():
tensor = tensor.contiguous()
return tensor.view(torch.uint8)
def _ensure_inputs(a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
if not (a.is_cuda and b.is_cuda and c.is_cuda):
raise RuntimeError("All tensors must live on CUDA for this kernel")
if c.size(1) != 1:
raise ValueError("Only single-column GEMV outputs are supported")
def custom_kernel(data: input_t) -> output_t:
a, b, _sfa, _sfb, sfa_perm, sfb_perm, c = data
_ensure_inputs(a, b, c)
device = a.device
module = _load_module()
profile = os.environ.get(_PROFILE_ENV) == "1"
if profile:
torch.cuda.synchronize(device)
evt_total_start = torch.cuda.Event(enable_timing=True)
evt_after_prep = torch.cuda.Event(enable_timing=True)
evt_after_launch = torch.cuda.Event(enable_timing=True)
evt_total_stop = torch.cuda.Event(enable_timing=True)
evt_total_start.record()
a_bytes = _prepare_matrix_bytes(a)
b_bytes = _prepare_vector_bytes(b)
sfa_perm_bytes = _prepare_perm_bytes(sfa_perm, device)
sfb_perm_bytes = _prepare_perm_bytes(sfb_perm, device)
if profile:
evt_after_prep.record()
m, k_bytes, batches = a_bytes.shape
if k_bytes % _BYTES_PER_BLOCK != 0:
raise ValueError("Packed K dimension must align to 16")
if b_bytes.shape[0] != k_bytes or b_bytes.shape[1] != batches:
raise ValueError("Vector layout does not match matrix layout")
if sfa_perm_bytes.dim() != 6 or sfb_perm_bytes.dim() != 6:
raise ValueError("Permuted scale tensors must be 6D")
if sfa_perm_bytes.size(-1) != batches or sfb_perm_bytes.size(-1) != batches:
raise ValueError("Scale tensors batch size mismatch")
out_view = c.select(1, 0)
module.nvfp4_gemv(
a_bytes,
b_bytes,
sfa_perm_bytes,
sfb_perm_bytes,
out_view,
)
if profile:
evt_after_launch.record()
evt_total_stop.record()
torch.cuda.synchronize(device)
prep_ms = evt_total_start.elapsed_time(evt_after_prep)
launch_ms = evt_after_prep.elapsed_time(evt_after_launch)
total_ms = evt_total_start.elapsed_time(evt_total_stop)
raise RuntimeError(
f"NVFP4_PROFILE prep_ms={prep_ms:.3f} launch_ms={launch_ms:.3f} total_ms={total_ms:.3f}"
)
return c
scrolls · 553 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