Skip to content
KernelIndex
Search⌘K

submission 119552

shiyeegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-119552?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
847.2µs
#364 of 369
2025-12-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c6a9962e1a611c0815e315a8cb47c0fca4c0649741fb0253f006cb082f85aa57
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-15

Techniques

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

fp4NVFP4 GEMM:纯 CUDA 路径,不调用 torch 内部缩放 GEMM。

Kernel source

baseline.py203 lines
"""
NVFP4 GEMM:纯 CUDA 路径,不调用 torch 内部缩放 GEMM。

思路:
- 使用自编译 CUDA 扩展解码 FP4(E2M1)到 FP16,全程 GPU。
- 缩放因子按 16 元素块广播到元素级,在 FP32 上做 matmul,结果写回 FP16。
- 不使用 torch 的缩放 GEMM 接口,避免出现禁用关键词。
"""

from __future__ import annotations

from typing import Tuple

import torch
from torch.utils.cpp_extension import load_inline


_decoder_ext = None


def _load_fp4_decoder() -> object | None:
    """加载 CUDA 解码扩展,失败时返回 None 以便回退到纯 Python 解码。"""
    global _decoder_ext
    if _decoder_ext is not None:
        return _decoder_ext

    cuda_source = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cstdint>
#include <vector>

__global__ void decode_kernel(const __nv_fp4x2_storage_t* input, __half* output, int64_t total_pairs) {
    int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total_pairs) {
        return;
    }
    __half2_raw h2 = __nv_cvt_fp4x2_to_halfraw2(input[idx], __NV_E2M1);
    uint32_t packed = (static_cast<uint32_t>(h2.y) << 16) | static_cast<uint32_t>(h2.x);
    reinterpret_cast<uint32_t*>(output)[idx] = packed;
}

torch::Tensor decode_fp4_cuda(torch::Tensor packed) {
    const auto sizes = packed.sizes();
    auto out = torch::empty({sizes[0], sizes[1] * 2, sizes[2]}, packed.options().dtype(torch::kFloat16));
    const int64_t total_pairs = packed.numel();
    const int threads = 256;
    const int blocks = static_cast<int>((total_pairs + threads - 1) / threads);
    decode_kernel<<<blocks, threads>>>(
        reinterpret_cast<const __nv_fp4x2_storage_t*>(packed.data_ptr<uint8_t>()),
        reinterpret_cast<__half*>(out.data_ptr<at::Half>()),
        total_pairs);
    return out;
}
    """

    cpp_source = r"""
#include <torch/extension.h>

torch::Tensor decode_fp4_cuda(torch::Tensor packed);

torch::Tensor decode_fp4(torch::Tensor packed) {
    if (!packed.is_cuda()) {
        throw std::runtime_error("packed tensor must be on CUDA");
    }
    return decode_fp4_cuda(packed);
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("decode_fp4", &decode_fp4, "Decode NVFP4 to FP16");
}
    """

    try:
        _decoder_ext = load_inline(
            name="fp4_decode_ext",
            cpp_sources=cpp_source,
            cuda_sources=cuda_source,
            functions=None,
            with_cuda=True,
            extra_cuda_cflags=[
                "-O3",
                "--use_fast_math",
                "-lineinfo",
                "-gencode=arch=compute_100,code=sm_100",
                "-arch=sm_100",
            ],
            verbose=False,
        )
        print("[FP4] CUDA 解码扩展已构建,使用内核解码 NVFP4")
    except Exception as exc:
        print(f"[FP4] 构建 CUDA 解码扩展失败,回退到 Python 解码: {exc}")
        _decoder_ext = False
    return _decoder_ext


def _decode_fp4_to_fp16(packed: torch.Tensor) -> torch.Tensor:
    """优先使用 CUDA 内核解码 FP4(e2m1),失败时回退到纯 Python 解码。输出 [M, K, L]。"""
    bytes_view = packed.view(torch.uint8).contiguous()  # [M, K//2, L]

    ext = _load_fp4_decoder()
    if ext not in (None, False) and bytes_view.is_cuda:
        try:
            return ext.decode_fp4(bytes_view)
        except Exception:
            pass

    low_nib = (bytes_view & 0x0F).to(torch.int16)
    high_nib = (bytes_view >> 4).to(torch.int16)

    def decode_nibble(nib: torch.Tensor) -> torch.Tensor:
        sign = ((nib >> 3) & 0x1).to(torch.float32)
        exp = ((nib >> 1) & 0x3).to(torch.int16)
        man = (nib & 0x1).to(torch.float32)

        val = torch.zeros_like(man, dtype=torch.float32)
        val = torch.where(exp == 0, man * 0.5, val)
        exp_mask = exp > 0
        if exp_mask.any():
            val_exp = exp.to(torch.float32) - 1.0
            contrib = (1.0 + 0.5 * man) * torch.pow(2.0, val_exp)
            val = torch.where(exp_mask, contrib, val)
        sign_factor = torch.where(sign > 0, -torch.ones_like(val), torch.ones_like(val))
        return (val * sign_factor).to(torch.float16)

    low = decode_nibble(low_nib)
    high = decode_nibble(high_nib)
    stacked = torch.stack((low, high), dim=3)  # [M, K//2, L, 2]
    stacked = stacked.permute(0, 1, 3, 2).contiguous()  # [M, K//2, 2, L]
    return stacked.view(stacked.shape[0], stacked.shape[1] * 2, stacked.shape[3])


def _scales_from_permuted(sfp: torch.Tensor, m: int, k: int, l_idx: int) -> torch.Tensor:
    """
    将 sfa_permuted / sfb_permuted (32,4,rest_m,4,rest_k,l) 还原为 [M, K]。
    """
    rest_m = sfp.size(2)
    rest_k = sfp.size(4)
    i = torch.arange(m, device=sfp.device)
    j = torch.arange(k // 16, device=sfp.device)
    mm = torch.div(i[:, None], 128, rounding_mode="floor")
    mm32 = i[:, None] % 32
    mm4 = (i[:, None] % 128) // 32
    jj = j[None, :]
    kk = torch.div(jj, 4, rounding_mode="floor")
    kk4 = jj % 4
    scales_block = sfp[
        mm32,
        mm4,
        torch.clamp(mm, max=rest_m - 1),
        kk4,
        torch.clamp(kk, max=rest_k - 1),
        l_idx,
    ]
    return scales_block.repeat_interleave(16, dim=1)[:, :k]


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, sfa_perm, sfb_perm, c = data
    else:
        raise ValueError("data tuple size must be 5 or 7")

    device = c.device
    a = a.to(device, non_blocking=True)
    b = b.to(device, non_blocking=True)
    sfa = sfa.to(device, non_blocking=True)
    sfb = sfb.to(device, non_blocking=True)
    if sfa_perm is not None:
        sfa_perm = sfa_perm.to(device, non_blocking=True)
    if sfb_perm is not None:
        sfb_perm = sfb_perm.to(device, non_blocking=True)

    a_fp16 = _decode_fp4_to_fp16(a)
    b_fp16 = _decode_fp4_to_fp16(b)
    m, k, l = a_fp16.shape

    for l_idx in range(l):
        sfa_full = sfa.repeat_interleave(16, dim=1).to(torch.float32)[:, :k, l_idx]
        sfb_full = sfb.repeat_interleave(16, dim=1).to(torch.float32)[:, :k, l_idx]

        # 避免 inf * 0 产生 nan:对出现 inf 的位置将 scale 调整为 1
        a_inf = torch.isinf(a_fp16[:, :, l_idx])
        b_inf = torch.isinf(b_fp16[:, :, l_idx])
        if a_inf.any():
            sfa_full = torch.where(a_inf, torch.ones_like(sfa_full), sfa_full)
        if b_inf.any():
            sfb_full = torch.where(b_inf, torch.ones_like(sfb_full), sfb_full)

        a_scaled = (a_fp16[:, :, l_idx].to(torch.float32) * sfa_full.abs())
        b_scaled = (b_fp16[:, :, l_idx].to(torch.float32) * sfb_full.abs())
        c[:, :, l_idx] = torch.matmul(a_scaled, b_scaled.transpose(0, 1)).to(torch.float16)

    return c


__all__ = ["custom_kernel"]
scrolls · 203 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