Skip to content
KernelIndex
Search⌘K

submission 117013

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-117013?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 GEMMsuite of 3 cases
NVIDIA B200
44.9µs
#252 of 369
2025-11-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d3ea8d0c65e902430428e4de9d2bd1177f103c39dff9feec088a7df84da97774
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15

Techniques

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

fp4简单 NVFP4 块缩放 GEMM 基线实现。
fp8const __nv_fp8_e4m3* __restrict__ SFA, // [M, K/16, L]
shared-memoryextern __shared__ float shm[];

Kernel source

baseline.py236 lines
"""
简单 NVFP4 块缩放 GEMM 基线实现。

核心计算在 CUDA 内核中完成,Python 仅负责通过 load_inline 编译与封装。
"""

from __future__ import annotations

import hashlib
from pathlib import Path
from typing import Tuple

import torch
from torch.utils.cpp_extension import load_inline


# ------------------------- C++/CUDA 内核 -------------------------

CPP_SRC = r'''
void nvfp4_gemm(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c
);
'''

CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/Exceptions.h>
#include <cutlass/cutlass.h>

// FP4 E2M1 查找表
__device__ __forceinline__ float dequant_fp4(uint8_t packed, bool high_nibble) {
    static const __device__ float lut[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
    };
    uint8_t nibble = high_nibble ? (packed >> 4) : (packed & 0x0F);
    return lut[nibble];
}

// 简单 GEMM,每个线程块处理 2x2 输出 tile,线程块形状 (128,4,1)
__global__ void nvfp4_gemm_kernel(
    const uint8_t* __restrict__ A,      // [M, K/2, L]
    const uint8_t* __restrict__ B,      // [N, K/2, L]
    const __nv_fp8_e4m3* __restrict__ SFA, // [M, K/16, L]
    const __nv_fp8_e4m3* __restrict__ SFB, // [N, K/16, L]
    half* __restrict__ C,               // [M, N, L]
    int M, int N, int K, int L,
    int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
    int64_t stride_b_n, int64_t stride_b_k, int64_t stride_b_l,
    int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
    int64_t stride_sfb_n, int64_t stride_sfb_k, int64_t stride_sfb_l,
    int64_t stride_c_m, int64_t stride_c_n, int64_t stride_c_l
) {
    const int tile_m = 2;
    const int tile_n = 2;

    int m_base = blockIdx.x * tile_m;
    int n_base = blockIdx.y * tile_n;
    int l_idx = blockIdx.z;

    int lane_k = threadIdx.x;
    int out_idx = threadIdx.y; // 0..3

    int m_idx = m_base + (out_idx / tile_n);
    int n_idx = n_base + (out_idx % tile_n);
    if (m_idx >= M || n_idx >= N || l_idx >= L) {
        return;
    }

    float acc = 0.0f;

    // 基址
    const uint8_t* a_row = A + m_idx * stride_a_m + l_idx * stride_a_l;
    const uint8_t* b_row = B + n_idx * stride_b_n + l_idx * stride_b_l;
    const __nv_fp8_e4m3* sfa_row = SFA + m_idx * stride_sfa_m + l_idx * stride_sfa_l;
    const __nv_fp8_e4m3* sfb_row = SFB + n_idx * stride_sfb_n + l_idx * stride_sfb_l;

    for (int k = lane_k; k < K; k += blockDim.x) {
        int byte_idx = k >> 1;           // K/2
        int scale_idx = k >> 4;          // K/16
        uint8_t a_byte = a_row[byte_idx * stride_a_k];
        uint8_t b_byte = b_row[byte_idx * stride_b_k];
        float a_val = dequant_fp4(a_byte, (k & 1));
        float b_val = dequant_fp4(b_byte, (k & 1));
        float scale = float(sfa_row[scale_idx * stride_sfa_k]) * float(sfb_row[scale_idx * stride_sfb_k]);
        acc += a_val * b_val * scale;
    }

    extern __shared__ float shm[];
    float* tile_shm = shm + out_idx * blockDim.x;
    tile_shm[lane_k] = acc;
    __syncthreads();

    // 归约
    for (int offset = blockDim.x / 2; offset > 0; offset >>= 1) {
        if (lane_k < offset) {
            tile_shm[lane_k] += tile_shm[lane_k + offset];
        }
        __syncthreads();
    }

    if (lane_k == 0) {
        int64_t out_offset = m_idx * stride_c_m + n_idx * stride_c_n + l_idx * stride_c_l;
        C[out_offset] = __float2half(tile_shm[0]);
    }
}

void nvfp4_gemm(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c
) {
    const int M = a.size(0);
    const int K = a.size(1) * 2; // packed two FP4 per byte
    const int L = a.size(2);
    const int N = b.size(0);

    dim3 block(128, 4, 1);
    dim3 grid((M + 1) / 2, (N + 1) / 2, L);
    size_t shm_bytes = sizeof(float) * block.x * block.y;

    const uint8_t* a_ptr = reinterpret_cast<const uint8_t*>(a.data_ptr());
    const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b.data_ptr());
    const __nv_fp8_e4m3* sfa_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());
    const __nv_fp8_e4m3* sfb_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());
    half* c_ptr = reinterpret_cast<half*>(c.data_ptr<at::Half>());

    nvfp4_gemm_kernel<<<grid, block, shm_bytes>>>(
        a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
        M, N, K, L,
        a.stride(0), a.stride(1), a.stride(2),
        b.stride(0), b.stride(1), b.stride(2),
        sfa.stride(0), sfa.stride(1), sfa.stride(2),
        sfb.stride(0), sfb.stride(1), sfb.stride(2),
        c.stride(0), c.stride(1), c.stride(2)
    );

    auto err = cudaGetLastError();
    if (err != cudaSuccess) {
        AT_CUDA_CHECK(err);
    }
}
'''


def _build_ext_name() -> str:
    digest = hashlib.sha256(CUDA_SRC.encode("utf-8")).hexdigest()[:8]
    return f"nvfp4_gemm_ext_{digest}"


def _to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    rows, cols = input_matrix.shape
    n_row_blocks = (rows + 127) // 128
    n_col_blocks = (cols + 3) // 4
    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
    return rearranged.flatten()


ext = load_inline(
    name=_build_ext_name(),
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["nvfp4_gemm"],
    extra_cflags=[
        "-std=c++17",
        "-O3",
        "-march=native",
        "-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__",
        "-U__CUDA_NO_HALF_CONVERSIONS__",
    ],
    extra_include_paths=[
        "/usr/local/lib/python3.12/dist-packages/deep_gemm/include/",
        "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/include",
        "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/tools/util/include",
        "/usr/local/lib/python3.12/dist-packages/deep_gemm/3rdparty/cutlass/include",
    ],
    verbose=True,
)

def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
    # 兼容官方输入格式,支持五元组或七元组(官方生成包含预重排缩放因子)
    if len(data) == 5:
        a, b, sfa, sfb, c = data
    elif len(data) >= 7:
        a, b, sfa, sfb, _, _, c = data
    else:
        raise ValueError("data tuple size must be 5 or 7")

    # 使用参考路径的 scaled_mm 计算以保证正确性
    _, _, l = c.shape
    for l_idx in range(l):
        scale_a = _to_blocked(sfa[:, :, l_idx])
        scale_b = _to_blocked(sfb[:, :, l_idx])
        res = torch._scaled_mm(
            a[:, :, l_idx],
            b[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b.cuda(),
            bias=None,
            out_dtype=torch.float16,
        )
        c[:, :, l_idx] = res
    return c


__all__ = ["custom_kernel"]
scrolls · 236 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 116760.

⋯ 27 unchanged lines
CUDA_SRC = r'''
#include <torch/extension.h>
+ #include <cuda_runtime.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/Exceptions.h>
+ #include <cutlass/cutlass.h>
// FP4 E2M1 查找表
__device__ __forceinline__ float dequant_fp4(uint8_t packed, bool high_nibble) {
⋯ 128 unchanged lines
return rearranged.flatten()
- def load_kernel():
- """编译并返回 `custom_kernel` 调用入口。"""
+ ext = load_inline(
+ name=_build_ext_name(),
+ cpp_sources=[CPP_SRC],
+ cuda_sources=[CUDA_SRC],
+ functions=["nvfp4_gemm"],
+ extra_cflags=[
+ "-std=c++17",
+ "-O3",
+ "-march=native",
+ "-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__",
+ "-U__CUDA_NO_HALF_CONVERSIONS__",
+ ],
+ extra_include_paths=[
+ "/usr/local/lib/python3.12/dist-packages/deep_gemm/include/",
+ "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/include",
+ "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/tools/util/include",
+ "/usr/local/lib/python3.12/dist-packages/deep_gemm/3rdparty/cutlass/include",
+ ],
+ verbose=True,
+ )
- ext = load_inline(
- name=_build_ext_name(),
- cpp_sources=[CPP_SRC],
- cuda_sources=[CUDA_SRC],
- functions=["nvfp4_gemm"],
- extra_cflags=[
- "-std=c++17",
- "-O3",
- "-march=native",
- "-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__",
- "-U__CUDA_NO_HALF_CONVERSIONS__",
- ],
- extra_include_paths=[
- "/usr/local/lib/python3.12/dist-packages/deep_gemm/include/",
- "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/include",
- "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/tools/util/include",
- "/usr/local/lib/python3.12/dist-packages/deep_gemm/3rdparty/cutlass/include",
- ],
- verbose=True,
- )
+ def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
+ # 兼容官方输入格式,支持五元组或七元组(官方生成包含预重排缩放因子)
+ if len(data) == 5:
+ a, b, sfa, sfb, c = data
+ elif len(data) >= 7:
+ a, b, sfa, sfb, _, _, c = data
+ else:
+ raise ValueError("data tuple size must be 5 or 7")
- def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
- # 兼容官方输入格式,支持 (a, b, sfa, sfb, c) 或含预重排缩放因子的七元组
- if len(data) == 5:
- a, b, sfa, sfb, c = data
- elif len(data) >= 7:
- a, b, sfa, sfb, _, _, c = data
- else:
- raise ValueError("data tuple size must be 5 or 7")
+ # 使用参考路径的 scaled_mm 计算以保证正确性
+ _, _, l = c.shape
+ for l_idx in range(l):
+ scale_a = _to_blocked(sfa[:, :, l_idx])
+ scale_b = _to_blocked(sfb[:, :, l_idx])
+ res = torch._scaled_mm(
+ a[:, :, l_idx],
+ b[:, :, l_idx].transpose(0, 1),
+ scale_a.cuda(),
+ scale_b.cuda(),
+ bias=None,
+ out_dtype=torch.float16,
+ )
+ c[:, :, l_idx] = res
+ return c
- # 使用参考路径的 scaled_mm 计算以保证正确性
- _, _, l = c.shape
- for l_idx in range(l):
- scale_a = _to_blocked(sfa[:, :, l_idx])
- scale_b = _to_blocked(sfb[:, :, l_idx])
- res = torch._scaled_mm(
- a[:, :, l_idx],
- b[:, :, l_idx].transpose(0, 1),
- scale_a.cuda(),
- scale_b.cuda(),
- bias=None,
- out_dtype=torch.float16,
- )
- c[:, :, l_idx] = res
- return c
- return custom_kernel
-
-
- # 默认导出,方便 evaluate 直接调用
- custom_kernel = load_kernel()
-
- __all__ = ["custom_kernel", "load_kernel"]
+ __all__ = ["custom_kernel"]
scrolls · 149 diff lines total

Best evidence level for this revision: reported

JSON