Skip to content
KernelIndex
Search⌘K

submission 70464

Infatoshi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 261 lines, June 9 Researcher Reciprocity License v1.0.

submission_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70464?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
738.5µs
#602 of 678
2025-11-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ea31c4d24904509d780466d66acf18ae2c3b09a0422b24a1501b00a6d6801a11
license declaredunknown
license concludedunknown
authorsInfatoshi
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4m.def("nvfp4_gemv", &nvfp4_gemv, "NVFP4 batched GEMV");

Kernel source

submission_v1.py261 lines
from functools import lru_cache

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_fp16.h>
#include <cstdint>
#include <cmath>

namespace {

constexpr int WARP_SIZE = 32;
constexpr int ROWS_PER_BLOCK = 4;

__device__ __forceinline__ float decode_fp4_e2m1(uint8_t v) {
  int sign = (v >> 3) & 0x1;
  int exp = (v >> 1) & 0x3;
  int mant = v & 0x1;

  float value;
  if (exp == 0) {
    if (mant == 0) {
      value = 0.0f;
    } else {
      float frac = float(mant) * 0.5f;
      value = ldexpf(frac, 0);  // 2^(1-bias) with bias=1
    }
  } else {
    float frac = 1.0f + float(mant) * 0.5f;
    value = ldexpf(frac, exp - 1);
  }
  return sign ? -value : value;
}

__device__ __forceinline__ float decode_fp8_e4m3(uint8_t v) {
  int sign = (v >> 7) & 0x1;
  int exp = (v >> 3) & 0xF;
  int mant = v & 0x7;

  const int bias = 7;
  float value;

  if (exp == 0) {
    if (mant == 0) {
      value = 0.0f;
    } else {
      float frac = float(mant) / 8.0f;
      value = ldexpf(frac, 1 - bias);
    }
  } else if (exp == 0xF) {
    // Finite-number saturate: clamp to the largest finite value representable.
    float frac = 1.0f + 7.0f / 8.0f;
    value = ldexpf(frac, (0xE - bias));
    value = ldexpf(value, 1); // multiply by 2 once more to match exp==14, mant==7
  } else {
    float frac = 1.0f + float(mant) / 8.0f;
    value = ldexpf(frac, exp - bias);
  }
  return sign ? -value : value;
}

struct Strides3 {
  int64_t d0;
  int64_t d1;
  int64_t d2;
};

__global__ void nvfp4_gemv_kernel(
    const uint8_t* __restrict__ A,
    Strides3 a_stride,
    const uint8_t* __restrict__ B,
    Strides3 b_stride,
    const uint8_t* __restrict__ SFA,
    Strides3 sfa_stride,
    const uint8_t* __restrict__ SFB,
    Strides3 sfb_stride,
    half* __restrict__ Out,
    Strides3 out_stride,
    int M,
    int K,
    int L,
    int packs_per_row,
    int blocks_per_row) {

  int warp_id = threadIdx.x / WARP_SIZE;
  int lane_id = threadIdx.x % WARP_SIZE;
  int row = blockIdx.x * ROWS_PER_BLOCK + warp_id;
  int batch = blockIdx.z;

  if (row >= M || batch >= L) {
    return;
  }

  float accum = 0.0f;

  for (int pack_idx = lane_id; pack_idx < packs_per_row; pack_idx += WARP_SIZE) {
    // Load packed nvfp4 data for A and B (first N row only).
    int64_t a_offset = row * a_stride.d0 + pack_idx * a_stride.d1 + batch * a_stride.d2;
    uint8_t packed_a = A[a_offset];

    int64_t b_offset = pack_idx * b_stride.d1 + batch * b_stride.d2;  // first row -> d0 term omitted
    uint8_t packed_b = B[b_offset];

    int col0 = pack_idx * 2;
    int block0 = col0 >> 4;

    int64_t sfa_off0 = row * sfa_stride.d0 + block0 * sfa_stride.d1 + batch * sfa_stride.d2;
    int64_t sfb_off0 = block0 * sfb_stride.d1 + batch * sfb_stride.d2;
    float scale_a0 = decode_fp8_e4m3(SFA[sfa_off0]);
    float scale_b0 = decode_fp8_e4m3(SFB[sfb_off0]);

    float base_a0 = decode_fp4_e2m1(packed_a & 0xF);
    float base_b0 = decode_fp4_e2m1(packed_b & 0xF);
    accum += (base_a0 * scale_a0) * (base_b0 * scale_b0);

    int col1 = col0 + 1;
    if (col1 < K) {
      int block1 = col1 >> 4;

      float scale_a1 = scale_a0;
      float scale_b1 = scale_b0;
      if (block1 != block0) {
        int64_t sfa_off1 = row * sfa_stride.d0 + block1 * sfa_stride.d1 + batch * sfa_stride.d2;
        int64_t sfb_off1 = block1 * sfb_stride.d1 + batch * sfb_stride.d2;
        scale_a1 = decode_fp8_e4m3(SFA[sfa_off1]);
        scale_b1 = decode_fp8_e4m3(SFB[sfb_off1]);
      }

      float base_a1 = decode_fp4_e2m1(packed_a >> 4);
      float base_b1 = decode_fp4_e2m1(packed_b >> 4);
      accum += (base_a1 * scale_a1) * (base_b1 * scale_b1);
    }
  }

  // Warp reduction
  for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
    accum += __shfl_down_sync(0xffffffff, accum, offset);
  }

  if (lane_id == 0) {
    int64_t out_offset = row * out_stride.d0 + batch * out_stride.d2;
    Out[out_offset] = __float2half(accum);
  }
}

}  // namespace

void nvfp4_gemv(torch::Tensor a,
                torch::Tensor b,
                torch::Tensor sfa,
                torch::Tensor sfb,
                torch::Tensor out) {
  TORCH_CHECK(a.is_cuda(), "A must be on CUDA");
  TORCH_CHECK(b.is_cuda(), "B must be on CUDA");
  TORCH_CHECK(sfa.is_cuda(), "SFA must be on CUDA");
  TORCH_CHECK(sfb.is_cuda(), "SFB must be on CUDA");
  TORCH_CHECK(out.is_cuda(), "Output must be on CUDA");

  TORCH_CHECK(a.dim() == 3, "A must be [M, K_pack, L]");
  TORCH_CHECK(b.dim() == 3, "B must be [N, K_pack, L]");
  TORCH_CHECK(out.dim() == 3, "Output must be [M, 1, L]");
  TORCH_CHECK(sfa.dim() == 3, "SFA must be [M, K_blk, L]");
  TORCH_CHECK(sfb.dim() == 3, "SFB must be [1, K_blk, L]");

  int64_t M = a.size(0);
  int64_t K = b.size(1) * 2;  // packs -> actual columns
  int64_t L = a.size(2);

  TORCH_CHECK(b.size(1) == a.size(1), "Packed K dimension mismatch");
  TORCH_CHECK(b.size(2) == L, "B batch mismatch");
  TORCH_CHECK(out.size(0) == M, "Output M mismatch");
  TORCH_CHECK(out.size(1) == 1, "Output N must be 1");
  TORCH_CHECK(out.size(2) == L, "Output batch mismatch");
  TORCH_CHECK(K % 16 == 0, "K must be divisible by 16");
  TORCH_CHECK(sfa.size(0) == M, "SFA M mismatch");
  TORCH_CHECK(sfa.size(1) * 16 >= K, "SFA K coverage mismatch");
  TORCH_CHECK(sfa.size(2) == L, "SFA batch mismatch");
  TORCH_CHECK(sfb.size(0) >= 1, "SFB requires at least one row");
  TORCH_CHECK(sfb.size(1) * 16 >= K, "SFB K coverage mismatch");
  TORCH_CHECK(sfb.size(2) == L, "SFB batch mismatch");

  auto stream = at::cuda::getCurrentCUDAStream();

  Strides3 a_stride{a.stride(0), a.stride(1), a.stride(2)};
  Strides3 b_stride{b.stride(0), b.stride(1), b.stride(2)};
  Strides3 sfa_stride{sfa.stride(0), sfa.stride(1), sfa.stride(2)};
  Strides3 sfb_stride{sfb.stride(0), sfb.stride(1), sfb.stride(2)};
  Strides3 out_stride{out.stride(0), out.stride(1), out.stride(2)};

  const uint8_t* ptr_a = reinterpret_cast<const uint8_t*>(a.data_ptr());
  const uint8_t* ptr_b = reinterpret_cast<const uint8_t*>(b.data_ptr());
  const uint8_t* ptr_sfa = reinterpret_cast<const uint8_t*>(sfa.data_ptr());
  const uint8_t* ptr_sfb = reinterpret_cast<const uint8_t*>(sfb.data_ptr());
  half* ptr_out = reinterpret_cast<half*>(out.data_ptr<at::Half>());

  int packs_per_row = static_cast<int>((K + 1) >> 1);
  int blocks_per_row = static_cast<int>((K + 15) >> 4);

  dim3 block_dim(WARP_SIZE * ROWS_PER_BLOCK);
  dim3 grid_dim((static_cast<int>(M) + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, static_cast<unsigned>(L));

  nvfp4_gemv_kernel<<<grid_dim, block_dim, 0, stream>>>(
      ptr_a, a_stride,
      ptr_b, b_stride,
      ptr_sfa, sfa_stride,
      ptr_sfb, sfb_stride,
      ptr_out, out_stride,
      static_cast<int>(M),
      static_cast<int>(K),
      static_cast<int>(L),
      packs_per_row,
      blocks_per_row);
  TORCH_CHECK(cudaGetLastError() == cudaSuccess, "nvfp4_gemv kernel launch failed");
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("nvfp4_gemv", &nvfp4_gemv, "NVFP4 batched GEMV");
}
"""


@lru_cache(maxsize=1)
def _load_module():
    extra_cuda_cflags = ["-O3", "--use_fast_math"]
    return load_inline(
        name="nvfp4_gemv_ext",
        cpp_sources="",
        cuda_sources=CUDA_SRC,
        extra_cuda_cflags=extra_cuda_cflags,
    )


def _ensure_device_tensor(tensor: torch.Tensor, device: torch.device) -> torch.Tensor:
    if tensor.device == device:
        return tensor
    return tensor.to(device=device, non_blocking=True)


def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, _sfa_perm, _sfb_perm, c = data
    device = a.device
    if device.type != "cuda":
        raise RuntimeError("nvfp4 kernel requires CUDA tensors")

    module = _load_module()

    a_dev = a
    b_dev = b
    sfa_dev = _ensure_device_tensor(sfa, device)
    sfb_dev = _ensure_device_tensor(sfb, device)
    c_dev = c

    module.nvfp4_gemv(a_dev, b_dev, sfa_dev, sfb_dev, c_dev)
    return c_dev
scrolls · 261 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