submission 476068
kaiming-cheng · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 195 lines, June 9 Researcher Reciprocity License v1.0.
test.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-476068?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:e575432190ca0c7cc4c1a8bb1d40d41f42fd773fae3605a5a3019746359f132c
license declaredunknown
license concludedunknown
authorskaiming-cheng
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""Optimized FP4 block-scaled GEMM kernel."""num-warps = 8
num_warps=8,stages = 4
num_stages=4,tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128Kernel source
test.py195 lines
import triton
import triton.language as tl
import torch
def ceil_div(a, b):
return (a + b - 1) // b
@triton.jit
def nvfp4_group_gemm_kernel_optimized(
# Pointers - a and b are uint8 packed FP4
a_ptr, b_ptr, c_ptr,
sfa_ptr, sfb_ptr,
# Dimensions
M, N, K,
# Strides for A [M, K//2] (packed uint8)
stride_am, stride_ak,
# Strides for B [N, K//2] (packed uint8)
stride_bn, stride_bk,
# Strides for C [M, N]
stride_cm, stride_cn,
# Scale strides - reordered format [32, 4, rest_m, 4, rest_k, L]
stride_sfa_0, stride_sfa_1, stride_sfa_2, stride_sfa_3, stride_sfa_4, stride_sfa_5,
stride_sfb_0, stride_sfb_1, stride_sfb_2, stride_sfb_3, stride_sfb_4, stride_sfb_5,
# L index for scale factors
l_idx,
# Block sizes
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
SF_VEC_SIZE: tl.constexpr,
):
"""Optimized FP4 block-scaled GEMM kernel."""
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
# Enhanced swizzle for better L2 cache locality
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
# Initialize accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Scale block size
BLOCK_K_SCALE: tl.constexpr = BLOCK_K // SF_VEC_SIZE
# Number of K iterations
num_k_iters = tl.cdiv(K, BLOCK_K)
# Base offsets for this tile - compute once
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k_packed = tl.arange(0, BLOCK_K // 2)
offs_scale_k = tl.arange(0, BLOCK_K_SCALE)
# Precompute scale indices for M dimension
mm32_a = offs_m % 32
mm4_a = (offs_m % 128) // 32
mm_a = offs_m // 128
# Precompute scale indices for N dimension
mm32_b = offs_n % 32
mm4_b = (offs_n % 128) // 32
mm_b = offs_n // 128
# K scale dimension
sf_k = tl.cdiv(K, SF_VEC_SIZE)
# Precompute base pointers
a_base = a_ptr + offs_m[:, None] * stride_am
b_base = b_ptr + offs_n[:, None] * stride_bn
# Precompute scale base pointers
sfa_base = (sfa_ptr +
mm32_a[:, None] * stride_sfa_0 +
mm4_a[:, None] * stride_sfa_1 +
mm_a[:, None] * stride_sfa_2 +
l_idx * stride_sfa_5)
sfb_base = (sfb_ptr +
mm32_b[:, None] * stride_sfb_0 +
mm4_b[:, None] * stride_sfb_1 +
mm_b[:, None] * stride_sfb_2 +
l_idx * stride_sfb_5)
# Masks for M and N (reused)
mask_m = offs_m < M
mask_n = offs_n < N
for k_iter in range(num_k_iters):
# K offset in packed format
k_offset_packed = k_iter * (BLOCK_K // 2)
# Load A tile
a_ptrs = a_base + (offs_k_packed[None, :] + k_offset_packed) * stride_ak
mask_a = mask_m[:, None] & ((offs_k_packed[None, :] + k_offset_packed) < (K // 2))
a = tl.load(a_ptrs, mask=mask_a, other=0)
# Load B tile
b_ptrs = b_base + (offs_k_packed[None, :] + k_offset_packed) * stride_bk
mask_b = mask_n[:, None] & ((offs_k_packed[None, :] + k_offset_packed) < (K // 2))
b = tl.load(b_ptrs, mask=mask_b, other=0)
# Scale factor K indices
scale_k_idx = k_iter * BLOCK_K_SCALE
col_idx = scale_k_idx + offs_scale_k
kk4 = col_idx % 4
kk = col_idx // 4
# Load scales for A
sfa_ptrs = sfa_base + kk4[None, :] * stride_sfa_3 + kk[None, :] * stride_sfa_4
mask_sfa = mask_m[:, None] & (col_idx[None, :] < sf_k)
scale_a = tl.load(sfa_ptrs, mask=mask_sfa, other=1.0)
# Load scales for B
sfb_ptrs = sfb_base + kk4[None, :] * stride_sfb_3 + kk[None, :] * stride_sfb_4
mask_sfb = mask_n[:, None] & (col_idx[None, :] < sf_k)
scale_b = tl.load(sfb_ptrs, mask=mask_sfb, other=1.0)
# Scaled dot product for FP4
acc = tl.dot_scaled(a, scale_a, "e2m1", tl.trans(b), scale_b, "e2m1", acc)
# Store output
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
mask_c = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
c = acc.to(tl.float16)
tl.store(c_ptrs, c, mask=mask_c)
def kernel_function(abc_tensors, sfasfb_reordered_tensors, problem_sizes):
"""NVFP4 block-scaled group GEMM wrapper using optimized Triton kernel."""
SF_VEC_SIZE = 16
results = []
for group_idx, ((a, b, c), (sfa_reordered, sfb_reordered), (m, n, k, l)) in enumerate(
zip(abc_tensors, sfasfb_reordered_tensors, problem_sizes)
):
a_uint8 = a.view(torch.uint8)
b_uint8 = b.view(torch.uint8)
for l_idx in range(l):
a_slice = a_uint8[:, :, l_idx].contiguous()
b_slice = b_uint8[:, :, l_idx].contiguous()
c_slice = c[:, :, l_idx]
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
GROUP_SIZE_M = 8
grid = (ceil_div(m, BLOCK_M) * ceil_div(n, BLOCK_N),)
nvfp4_group_gemm_kernel_optimized[grid](
a_slice, b_slice, c_slice,
sfa_reordered, sfb_reordered,
m, n, k,
a_slice.stride(0), a_slice.stride(1),
b_slice.stride(0), b_slice.stride(1),
c_slice.stride(0), c_slice.stride(1),
sfa_reordered.stride(0), sfa_reordered.stride(1), sfa_reordered.stride(2),
sfa_reordered.stride(3), sfa_reordered.stride(4), sfa_reordered.stride(5),
sfb_reordered.stride(0), sfb_reordered.stride(1), sfb_reordered.stride(2),
sfb_reordered.stride(3), sfb_reordered.stride(4), sfb_reordered.stride(5),
l_idx,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
GROUP_SIZE_M=GROUP_SIZE_M,
SF_VEC_SIZE=SF_VEC_SIZE,
num_warps=8,
num_stages=4,
)
results.append(c)
return results
# Wrapper for the evaluation system which passes a single data argument
def custom_kernel(data):
"""Wrapper that unpacks data from the evaluation system."""
abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
return kernel_function(abc_tensors, sfasfb_reordered_tensors, problem_sizes)
scrolls · 195 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