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
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 annotationsfrom typing import Tupleimport 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