Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
521.3µs
#596 of 678
2025-11-11

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