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
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.
fp4
m.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