Skip to content
KernelIndex
Search⌘K

submission 121413

shiyeegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

baseline_scaled_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-121413?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
17.7µs
#182 of 369
2025-12-04

Reported · How evidence levels are derived →

Source and license

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

Kernel source

baseline_scaled_mm.py36 lines
from __future__ import annotations

from typing import Tuple

import torch


def _prepare_scales_batch(scale_perm: torch.Tensor) -> torch.Tensor:
    permuted = scale_perm.permute(5, 2, 4, 0, 1, 3)
    return permuted.contiguous().view(scale_perm.size(-1), -1)


@torch.inference_mode()
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
    a, b, _, _, sfa_perm, sfb_perm, c = data

    _, _, l = c.shape

    scale_a_batch = _prepare_scales_batch(sfa_perm)
    scale_b_batch = _prepare_scales_batch(sfb_perm)

    for i in range(l):
        res = torch._scaled_mm(
            a[:, :, i],
            b[:, :, i].transpose(0, 1),
            scale_a_batch[i],
            scale_b_batch[i],
            bias=None,
            out_dtype=torch.float16,
        )
        c[:, :, i].copy_(res)

    return c


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

- """
- 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 _prepare_scales_batch(scale_perm: torch.Tensor) -> torch.Tensor:
+ permuted = scale_perm.permute(5, 2, 4, 0, 1, 3)
+ return permuted.contiguous().view(scale_perm.size(-1), -1)
- def _load_fp4_decoder() -> object | None:
- """加载 CUDA 解码扩展,失败时返回 None 以便回退到纯 Python 解码。"""
- global _decoder_ext
- if _decoder_ext is not None:
- return _decoder_ext
+ @torch.inference_mode()
+ def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
+ a, b, _, _, sfa_perm, sfb_perm, c = data
- 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>
+ _, _, l = c.shape
- __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;
- }
+ scale_a_batch = _prepare_scales_batch(sfa_perm)
+ scale_b_batch = _prepare_scales_batch(sfb_perm)
- 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,
+ for i in range(l):
+ res = torch._scaled_mm(
+ a[:, :, i],
+ b[:, :, i].transpose(0, 1),
+ scale_a_batch[i],
+ scale_b_batch[i],
+ bias=None,
+ out_dtype=torch.float16,
)
- print("[FP4] CUDA 解码扩展已构建,使用内核解码 NVFP4")
- except Exception as exc:
- print(f"[FP4] 构建 CUDA 解码扩展失败,回退到 Python 解码: {exc}")
- _decoder_ext = False
- return _decoder_ext
+ c[:, :, i].copy_(res)
-
- 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"]
+ __all__ = ["custom_kernel"]
No newline at end of file
scrolls · 222 diff lines total

Best evidence level for this revision: reported

JSON