submission 117310
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 77 lines, June 9 Researcher Reciprocity License v1.0.
node_9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-117310?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:8dd01ec58f465f29b2d70918df9cc3d746e8c460a4baf92e51e0ee48fc98a178
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
使用 PyTorch 内置 `torch._scaled_mm` 完成 NVFP4 块缩放 GEMM。Kernel source
node_9.py77 lines
"""
使用 PyTorch 内置 `torch._scaled_mm` 完成 NVFP4 块缩放 GEMM。
优先利用评测侧提供的预重排缩放因子,减少 Python 端重排开销;若未提供则退回参考重排。
"""
from __future__ import annotations
from typing import Tuple
import torch
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
def _to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
"""将缩放因子转换为 torch._scaled_mm 期望的分块布局。"""
rows, cols = input_matrix.shape
n_row_blocks = _ceil_div(rows, 128)
n_col_blocks = _ceil_div(cols, 4)
blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def _permuted_to_blocked(scale_permuted: torch.Tensor, l_idx: int) -> torch.Tensor:
"""
将评测侧预重排的缩放因子恢复到 torch._scaled_mm 可接受的扁平布局。
预重排形状约为 [32, 4, ceil(m/128), 4, ceil(k/16/4), L]。
"""
# 先取出指定 batch,再调整维度顺序使得 block_m、block_k 成为前两维,确保与参考重排一致。
sliced = scale_permuted[..., l_idx] # (32, 4, block_m, 4, block_k)
blocked = sliced.permute(2, 4, 0, 1, 3).contiguous().reshape(-1, 32, 16)
return blocked.flatten()
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
"""
兼容五元组 (a, b, sfa, sfb, c) 与七元组 (a, b, sfa, sfb, sfa_perm, sfb_perm, c)。
优先使用预重排缩放因子以减少重排成本。
"""
if len(data) == 5:
a, b, sfa, sfb, c = data
sfa_perm = sfb_perm = None
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")
_, _, l = c.shape
for l_idx in range(l):
# 缩放优先走预重排路径,缺失时回退参考重排。
if sfa_perm is not None and sfa_perm.dim() == 6:
scale_a = _permuted_to_blocked(sfa_perm, l_idx)
else:
scale_a = _to_blocked(sfa[:, :, l_idx])
if sfb_perm is not None and sfb_perm.dim() == 6:
scale_b = _permuted_to_blocked(sfb_perm, l_idx)
else:
scale_b = _to_blocked(sfb[:, :, l_idx])
result = torch._scaled_mm(
a[:, :, l_idx],
b[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b,
bias=None,
out_dtype=torch.float16,
)
c[:, :, l_idx].copy_(result)
return c
__all__ = ["custom_kernel"]
scrolls · 77 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 117013.
"""- 简单 NVFP4 块缩放 GEMM 基线实现。-- 核心计算在 CUDA 内核中完成,Python 仅负责通过 load_inline 编译与封装。+ 使用 PyTorch 内置 `torch._scaled_mm` 完成 NVFP4 块缩放 GEMM。+ 优先利用评测侧提供的预重排缩放因子,减少 Python 端重排开销;若未提供则退回参考重排。"""from __future__ import annotations- import hashlib- from pathlib import Pathfrom typing import Tupleimport torch- from torch.utils.cpp_extension import load_inline- # ------------------------- C++/CUDA 内核 -------------------------+ def _ceil_div(a: int, b: int) -> int:+ return (a + b - 1) // b- CPP_SRC = r'''- void nvfp4_gemm(- torch::Tensor a,- torch::Tensor b,- torch::Tensor sfa,- torch::Tensor sfb,- torch::Tensor c- );- '''- CUDA_SRC = r'''- #include <torch/extension.h>- #include <cuda_runtime.h>- #include <cuda.h>- #include <cuda_fp16.h>- #include <cuda_fp4.h>- #include <cuda_fp8.h>- #include <ATen/cuda/CUDAContext.h>- #include <ATen/cuda/Exceptions.h>- #include <cutlass/cutlass.h>-- // FP4 E2M1 查找表- __device__ __forceinline__ float dequant_fp4(uint8_t packed, bool high_nibble) {- static const __device__ float lut[16] = {- 0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,- -0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f- };- uint8_t nibble = high_nibble ? (packed >> 4) : (packed & 0x0F);- return lut[nibble];- }-- // 简单 GEMM,每个线程块处理 2x2 输出 tile,线程块形状 (128,4,1)- __global__ void nvfp4_gemm_kernel(- const uint8_t* __restrict__ A, // [M, K/2, L]- const uint8_t* __restrict__ B, // [N, K/2, L]- const __nv_fp8_e4m3* __restrict__ SFA, // [M, K/16, L]- const __nv_fp8_e4m3* __restrict__ SFB, // [N, K/16, L]- half* __restrict__ C, // [M, N, L]- int M, int N, int K, int L,- int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,- int64_t stride_b_n, int64_t stride_b_k, int64_t stride_b_l,- int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,- int64_t stride_sfb_n, int64_t stride_sfb_k, int64_t stride_sfb_l,- int64_t stride_c_m, int64_t stride_c_n, int64_t stride_c_l- ) {- const int tile_m = 2;- const int tile_n = 2;-- int m_base = blockIdx.x * tile_m;- int n_base = blockIdx.y * tile_n;- int l_idx = blockIdx.z;-- int lane_k = threadIdx.x;- int out_idx = threadIdx.y; // 0..3-- int m_idx = m_base + (out_idx / tile_n);- int n_idx = n_base + (out_idx % tile_n);- if (m_idx >= M || n_idx >= N || l_idx >= L) {- return;- }-- float acc = 0.0f;-- // 基址- const uint8_t* a_row = A + m_idx * stride_a_m + l_idx * stride_a_l;- const uint8_t* b_row = B + n_idx * stride_b_n + l_idx * stride_b_l;- const __nv_fp8_e4m3* sfa_row = SFA + m_idx * stride_sfa_m + l_idx * stride_sfa_l;- const __nv_fp8_e4m3* sfb_row = SFB + n_idx * stride_sfb_n + l_idx * stride_sfb_l;-- for (int k = lane_k; k < K; k += blockDim.x) {- int byte_idx = k >> 1; // K/2- int scale_idx = k >> 4; // K/16- uint8_t a_byte = a_row[byte_idx * stride_a_k];- uint8_t b_byte = b_row[byte_idx * stride_b_k];- float a_val = dequant_fp4(a_byte, (k & 1));- float b_val = dequant_fp4(b_byte, (k & 1));- float scale = float(sfa_row[scale_idx * stride_sfa_k]) * float(sfb_row[scale_idx * stride_sfb_k]);- acc += a_val * b_val * scale;- }-- extern __shared__ float shm[];- float* tile_shm = shm + out_idx * blockDim.x;- tile_shm[lane_k] = acc;- __syncthreads();-- // 归约- for (int offset = blockDim.x / 2; offset > 0; offset >>= 1) {- if (lane_k < offset) {- tile_shm[lane_k] += tile_shm[lane_k + offset];- }- __syncthreads();- }-- if (lane_k == 0) {- int64_t out_offset = m_idx * stride_c_m + n_idx * stride_c_n + l_idx * stride_c_l;- C[out_offset] = __float2half(tile_shm[0]);- }- }-- void nvfp4_gemm(- torch::Tensor a,- torch::Tensor b,- torch::Tensor sfa,- torch::Tensor sfb,- torch::Tensor c- ) {- const int M = a.size(0);- const int K = a.size(1) * 2; // packed two FP4 per byte- const int L = a.size(2);- const int N = b.size(0);-- dim3 block(128, 4, 1);- dim3 grid((M + 1) / 2, (N + 1) / 2, L);- size_t shm_bytes = sizeof(float) * block.x * block.y;-- const uint8_t* a_ptr = reinterpret_cast<const uint8_t*>(a.data_ptr());- const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b.data_ptr());- const __nv_fp8_e4m3* sfa_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());- const __nv_fp8_e4m3* sfb_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());- half* c_ptr = reinterpret_cast<half*>(c.data_ptr<at::Half>());-- nvfp4_gemm_kernel<<<grid, block, shm_bytes>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,- M, N, K, L,- a.stride(0), a.stride(1), a.stride(2),- b.stride(0), b.stride(1), b.stride(2),- sfa.stride(0), sfa.stride(1), sfa.stride(2),- sfb.stride(0), sfb.stride(1), sfb.stride(2),- c.stride(0), c.stride(1), c.stride(2)- );-- auto err = cudaGetLastError();- if (err != cudaSuccess) {- AT_CUDA_CHECK(err);- }- }- '''--- def _build_ext_name() -> str:- digest = hashlib.sha256(CUDA_SRC.encode("utf-8")).hexdigest()[:8]- return f"nvfp4_gemm_ext_{digest}"--def _to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:+ """将缩放因子转换为 torch._scaled_mm 期望的分块布局。"""rows, cols = input_matrix.shape- n_row_blocks = (rows + 127) // 128- n_col_blocks = (cols + 3) // 4- padded = input_matrix- blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)+ n_row_blocks = _ceil_div(rows, 128)+ n_col_blocks = _ceil_div(cols, 4)+ blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)return rearranged.flatten()- ext = load_inline(- name=_build_ext_name(),- cpp_sources=[CPP_SRC],- cuda_sources=[CUDA_SRC],- functions=["nvfp4_gemm"],- extra_cflags=[- "-std=c++17",- "-O3",- "-march=native",- "-fno-math-errno",- "-Wall",- ],- extra_cuda_cflags=[- "-O3",- "--use_fast_math",- "--extra-device-vectorization",- "--restrict",- "-std=c++17",- "--ptxas-options=-O3",- "--expt-relaxed-constexpr",- "-arch=sm_100a",- "-Xptxas",- "-v",- "-lineinfo",- "-U__CUDA_NO_HALF_OPERATORS__",- "-U__CUDA_NO_HALF_CONVERSIONS__",- ],- extra_include_paths=[- "/usr/local/lib/python3.12/dist-packages/deep_gemm/include/",- "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/include",- "/usr/local/lib/python3.12/dist-packages/flashinfer/data/cutlass/tools/util/include",- "/usr/local/lib/python3.12/dist-packages/deep_gemm/3rdparty/cutlass/include",- ],- verbose=True,- )+ def _permuted_to_blocked(scale_permuted: torch.Tensor, l_idx: int) -> torch.Tensor:+ """+ 将评测侧预重排的缩放因子恢复到 torch._scaled_mm 可接受的扁平布局。+ 预重排形状约为 [32, 4, ceil(m/128), 4, ceil(k/16/4), L]。+ """+ # 先取出指定 batch,再调整维度顺序使得 block_m、block_k 成为前两维,确保与参考重排一致。+ sliced = scale_permuted[..., l_idx] # (32, 4, block_m, 4, block_k)+ blocked = sliced.permute(2, 4, 0, 1, 3).contiguous().reshape(-1, 32, 16)+ return blocked.flatten()+def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:- # 兼容官方输入格式,支持五元组或七元组(官方生成包含预重排缩放因子)+ """+ 兼容五元组 (a, b, sfa, sfb, c) 与七元组 (a, b, sfa, sfb, sfa_perm, sfb_perm, c)。+ 优先使用预重排缩放因子以减少重排成本。+ """if len(data) == 5:a, b, sfa, sfb, c = data+ sfa_perm = sfb_perm = Noneelif len(data) >= 7:- a, b, sfa, sfb, _, _, c = data+ a, b, sfa, sfb, sfa_perm, sfb_perm, c = dataelse:raise ValueError("data tuple size must be 5 or 7")- # 使用参考路径的 scaled_mm 计算以保证正确性_, _, l = c.shapefor l_idx in range(l):- scale_a = _to_blocked(sfa[:, :, l_idx])- scale_b = _to_blocked(sfb[:, :, l_idx])- res = torch._scaled_mm(+ # 缩放优先走预重排路径,缺失时回退参考重排。+ if sfa_perm is not None and sfa_perm.dim() == 6:+ scale_a = _permuted_to_blocked(sfa_perm, l_idx)+ else:+ scale_a = _to_blocked(sfa[:, :, l_idx])++ if sfb_perm is not None and sfb_perm.dim() == 6:+ scale_b = _permuted_to_blocked(sfb_perm, l_idx)+ else:+ scale_b = _to_blocked(sfb[:, :, l_idx])++ result = torch._scaled_mm(a[:, :, l_idx],b[:, :, l_idx].transpose(0, 1),- scale_a.cuda(),- scale_b.cuda(),+ scale_a,+ scale_b,bias=None,out_dtype=torch.float16,)- c[:, :, l_idx] = res+ c[:, :, l_idx].copy_(result)return c
scrolls · 273 diff lines total
Best evidence level for this revision: reported
JSON