Skip to content
KernelIndex
Search⌘K

submission 465744

shigao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-465744?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 group GEMMsuite of 4 cases
NVIDIA B200
2.10ms
#145 of 145
2026-02-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9f786f124607cbea1fc4e137b5ddb83e3cb8823c73fdfc705f85c7982851c471
license declaredunknown
license concludedunknown
authorsshigao
imported2026-08-26

Techniques

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

fp4m.def("dequant_fp4_to_f16", &dequant_fp4_to_f16, "dequant fp4 to f16 (CUDA)");

Kernel source

submission.py262 lines
from __future__ import annotations

import os
from typing import List

import torch
from torch.utils.cpp_extension import load_inline

_EXT_MOD = None


def _get_ext_mod():
    global _EXT_MOD
    if _EXT_MOD is not None:
        return _EXT_MOD

    cuda_src = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <math.h>

// Vectorized FP4 to FP16 dequantization with scale
// Process 8 FP4 elements (4 bytes) per thread for better memory throughput

__device__ __forceinline__ float fp4_e2m1fn_to_f32(uint8_t v) {
    uint8_t sign = (v >> 3) & 1;
    uint8_t exp = (v >> 1) & 3;
    uint8_t man = v & 1;
    float mag;
    if (exp == 0) {
        mag = man ? 0.5f : 0.0f;
    } else {
        int e = int(exp) - 1;
        float base = ldexpf(1.0f, e);
        mag = base * (1.0f + 0.5f * float(man));
    }
    return sign ? -mag : mag;
}

// Optimized dequantization kernel with vectorized loads
__global__ void dequant_fp4_to_f16_kernel(
    const uint8_t* __restrict__ inp,
    const __half* __restrict__ scale,
    __half* __restrict__ out,
    int rows,
    int k,
    int l
) {
    // Each thread processes 8 elements (vectorized)
    const int vec_size = 8;
    int tid = blockIdx.x * blockDim.x + threadIdx.x;
    int total_elems = rows * k * l;
    int num_vec = total_elems / vec_size;
    
    int vec_idx = tid;
    if (vec_idx >= num_vec) {
        // Handle remaining elements
        int base_idx = tid + num_vec * vec_size;
        if (base_idx >= total_elems) return;
        
        int li = base_idx % l;
        int tmp = base_idx / l;
        int ki = tmp % k;
        int ri = tmp / k;
        
        int k_half = k >> 1;
        int k16 = k >> 4;
        int in_off = (ri * k_half + (ki >> 1)) * l + li;
        uint8_t packed = inp[in_off];
        uint8_t nib = (ki & 1) ? (packed >> 4) : (packed & 0x0F);
        
        float x = fp4_e2m1fn_to_f32(nib);
        float s = __half2float(scale[(ri * k16 + (ki >> 4)) * l + li]);
        out[base_idx] = __float2half_rn(x * s);
        return;
    }
    
    int base_idx = vec_idx * vec_size;
    
    #pragma unroll
    for (int i = 0; i < vec_size; ++i) {
        int idx = base_idx + i;
        int li = idx % l;
        int tmp = idx / l;
        int ki = tmp % k;
        int ri = tmp / k;
        
        int k_half = k >> 1;
        int k16 = k >> 4;
        int in_off = (ri * k_half + (ki >> 1)) * l + li;
        uint8_t packed = inp[in_off];
        uint8_t nib = (ki & 1) ? (packed >> 4) : (packed & 0x0F);
        
        float x = fp4_e2m1fn_to_f32(nib);
        float s = __half2float(scale[(ri * k16 + (ki >> 4)) * l + li]);
        out[idx] = __float2half_rn(x * s);
    }
}

// Simple coalesced dequantization - one thread per element
__global__ void dequant_fp4_coalesced_kernel(
    const uint8_t* __restrict__ inp,
    const __half* __restrict__ scale,
    __half* __restrict__ out,
    int rows,
    int k,
    int l
) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int total = rows * k * l;
    if (idx >= total) return;
    
    int li = idx % l;
    int tmp = idx / l;
    int ki = tmp % k;
    int ri = tmp / k;
    
    int k_half = k >> 1;
    int k16 = k >> 4;
    int in_off = (ri * k_half + (ki >> 1)) * l + li;
    uint8_t packed = inp[in_off];
    uint8_t nib = (ki & 1) ? (packed >> 4) : (packed & 0x0F);
    
    float x = fp4_e2m1fn_to_f32(nib);
    float s = __half2float(scale[(ri * k16 + (ki >> 4)) * l + li]);
    out[idx] = __float2half_rn(x * s);
}

static void dequant_fp4_to_f16(torch::Tensor inp, torch::Tensor scale, torch::Tensor out) {
    TORCH_CHECK(inp.is_cuda(), "inp must be CUDA");
    TORCH_CHECK(scale.is_cuda(), "scale must be CUDA");
    TORCH_CHECK(out.is_cuda(), "out must be CUDA");
    TORCH_CHECK(inp.is_contiguous(), "inp must be contiguous");
    TORCH_CHECK(scale.is_contiguous(), "scale must be contiguous");
    TORCH_CHECK(out.is_contiguous(), "out must be contiguous");
    TORCH_CHECK(scale.scalar_type() == at::kHalf, "scale must be float16");
    TORCH_CHECK(out.scalar_type() == at::kHalf, "out must be float16");
    TORCH_CHECK(inp.element_size() == 1, "inp element size must be 1 byte");
    TORCH_CHECK(inp.dim() == 3, "inp must be 3D");
    TORCH_CHECK(scale.dim() == 3, "scale must be 3D");
    TORCH_CHECK(out.dim() == 3, "out must be 3D");
    
    int rows = int(out.size(0));
    int k = int(out.size(1));
    int l = int(out.size(2));
    TORCH_CHECK(int(inp.size(0)) == rows, "rows mismatch");
    TORCH_CHECK(int(inp.size(2)) == l, "L mismatch");
    TORCH_CHECK(int(inp.size(1)) * 2 == k, "K mismatch for packed input");
    TORCH_CHECK(int(scale.size(0)) == rows, "scale rows mismatch");
    TORCH_CHECK(int(scale.size(2)) == l, "scale L mismatch");
    TORCH_CHECK(int(scale.size(1)) * 16 == k, "scale K mismatch (scale is per-16)");
    
    const uint8_t* inp_ptr = static_cast<const uint8_t*>(inp.data_ptr());
    const __half* scale_ptr = reinterpret_cast<const __half*>(scale.data_ptr<at::Half>());
    __half* out_ptr = reinterpret_cast<__half*>(out.data_ptr<at::Half>());
    
    int total = rows * k * l;
    int threads = 512;
    int blocks = (total + threads - 1) / threads;
    
    dequant_fp4_coalesced_kernel<<<blocks, threads>>>(inp_ptr, scale_ptr, out_ptr, rows, k, l);
    
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "kernel launch failed");
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("dequant_fp4_to_f16", &dequant_fp4_to_f16, "dequant fp4 to f16 (CUDA)");
}
"""

    extra_cuda = [
        "-O3",
        "-lineinfo",
        "-U__CUDA_NO_HALF_OPERATORS__",
        "-U__CUDA_NO_HALF_CONVERSIONS__",
        "--use_fast_math",
        "-gencode", "arch=compute_100a,code=sm_100a",
    ]
    extra_cxx = ["-O3"]

    build_dir = os.path.join(os.path.dirname(__file__), ".build_nvfp4_opt")
    os.makedirs(build_dir, exist_ok=True)
    _EXT_MOD = load_inline(
        name="nvfp4_group_gemm_opt_v2",
        cpp_sources="",
        cuda_sources=cuda_src,
        functions=None,
        extra_cuda_cflags=extra_cuda,
        extra_cflags=extra_cxx,
        with_cuda=True,
        build_directory=build_dir,
        verbose=False,
    )
    return _EXT_MOD


def _as_u8_view(x: torch.Tensor) -> torch.Tensor:
    if x.dtype == torch.uint8:
        return x
    if x.element_size() != 1:
        raise RuntimeError("packed dtype must have 1-byte elements")
    return x.view(torch.uint8)


def custom_kernel(data):
    abc_tensors, sfasfb_tensors, _sfasfb_reordered_tensors, problem_sizes = data
    g = len(problem_sizes)
    ext = _get_ext_mod()
    
    outs: List[torch.Tensor] = []
    
    for i in range(g):
        a, b, c = abc_tensors[i]
        sfa, sfb = sfasfb_tensors[i]
        m, n, k, l = problem_sizes[i]
        
        if not a.is_cuda:
            raise RuntimeError("only CUDA tensors are supported")
        
        
        sfa16 = sfa.to(device=a.device, dtype=torch.float16).contiguous()
        sfb16 = sfb.to(device=b.device, dtype=torch.float16).contiguous()
        
        
        a16 = torch.empty((int(m), int(k), int(l)), device=a.device, dtype=torch.float16)
        b16 = torch.empty((int(n), int(k), int(l)), device=b.device, dtype=torch.float16)
        
        
        ext.dequant_fp4_to_f16(_as_u8_view(a).contiguous(), sfa16, a16)
        ext.dequant_fp4_to_f16(_as_u8_view(b).contiguous(), sfb16, b16)
        
        
        
        
        
        c_out = c
        if not c_out.is_contiguous():
            c_tmp = torch.empty_like(c_out, memory_format=torch.contiguous_format)
        else:
            c_tmp = c_out
        
        
        for li in range(l):
            a_slice = a16[:, :, li]  
            b_slice = b16[:, :, li]  
            
            c_slice = torch.matmul(a_slice, b_slice.t())
            c_tmp[:, :, li] = c_slice
        
        if c_tmp is not c_out:
            c_out.copy_(c_tmp)
        outs.append(c_out)
    
    return outs


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