submission 387939
vesper15 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 294 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-387939?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:1d641005b5a30ed084089a460293b883b6282e4952bf6c59d457227ced864392
license declaredunknown
license concludedunknown
authorsvesper15
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
This implementation computes grouped matrix multiplication with FP4 inputsKernel source
submission.py294 lines
"""
Optimized Block-Scaled Group GEMM Kernel for NVIDIA B200 (Blackwell)
This implementation computes grouped matrix multiplication with FP4 inputs
and FP8 block scaling using torch._scaled_mm for optimal tensor core utilization.
"""
import torch
from torch.utils.cpp_extension import load_inline
# CUDA kernel for scale factor conversion to blocked format
cuda_source = """
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// Process both A and B scale factors in one kernel launch
__global__ void process_scales_kernel(
const uint8_t* __restrict__ sfa,
const uint8_t* __restrict__ sfb,
uint8_t* __restrict__ scale_a,
uint8_t* __restrict__ scale_b,
const int M, const int N,
const int K16,
const int n_row_blocks_a, const int n_col_blocks,
const int n_row_blocks_b
) {
const int tid = blockIdx.x * blockDim.x + threadIdx.x;
const int stride = blockDim.x * gridDim.x;
const int K16_vec = K16 >> 2;
const int K16_rem = K16 & 3;
const int base_k = K16_vec << 2;
if (K16_vec > 0) {
const int total_vec_a = M * K16_vec;
for (int vec = tid; vec < total_vec_a; vec += stride) {
const int m = vec / K16_vec;
const int k4 = vec % K16_vec;
const int row_block = m >> 7;
const int in_row = m & 127;
const int i32 = in_row & 31;
const int i4 = in_row >> 5;
const int block_idx = row_block * n_col_blocks + k4;
const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2);
const int base_in = m * K16 + (k4 << 2);
const uint32_t val = *reinterpret_cast<const uint32_t*>(sfa + base_in);
*reinterpret_cast<uint32_t*>(scale_a + out_idx) = val;
}
const int total_vec_b = N * K16_vec;
for (int vec = tid; vec < total_vec_b; vec += stride) {
const int n = vec / K16_vec;
const int k4 = vec % K16_vec;
const int row_block = n >> 7;
const int in_row = n & 127;
const int i32 = in_row & 31;
const int i4 = in_row >> 5;
const int block_idx = row_block * n_col_blocks + k4;
const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2);
const int base_in = n * K16 + (k4 << 2);
const uint32_t val = *reinterpret_cast<const uint32_t*>(sfb + base_in);
*reinterpret_cast<uint32_t*>(scale_b + out_idx) = val;
}
}
if (K16_rem > 0) {
const int total_rem_a = M * K16_rem;
for (int rem = tid; rem < total_rem_a; rem += stride) {
const int m = rem / K16_rem;
const int k_offset = rem % K16_rem;
const int k = base_k + k_offset;
if (k >= K16) continue;
const int row_block = m >> 7;
const int in_row = m & 127;
const int i32 = in_row & 31;
const int i4 = in_row >> 5;
const int col_block = k >> 2;
const int in_col = k & 3;
const int block_idx = row_block * n_col_blocks + col_block;
const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2) + in_col;
scale_a[out_idx] = sfa[m * K16 + k];
}
const int total_rem_b = N * K16_rem;
for (int rem = tid; rem < total_rem_b; rem += stride) {
const int n = rem / K16_rem;
const int k_offset = rem % K16_rem;
const int k = base_k + k_offset;
if (k >= K16) continue;
const int row_block = n >> 7;
const int in_row = n & 127;
const int i32 = in_row & 31;
const int i4 = in_row >> 5;
const int col_block = k >> 2;
const int in_col = k & 3;
const int block_idx = row_block * n_col_blocks + col_block;
const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2) + in_col;
scale_b[out_idx] = sfb[n * K16 + k];
}
}
}
std::vector<torch::Tensor> process_scales(
torch::Tensor sfa, torch::Tensor sfb
) {
const int M = sfa.size(0);
const int N = sfb.size(0);
const int K16 = sfa.size(1);
const int n_row_blocks_a = (M + 127) / 128;
const int n_row_blocks_b = (N + 127) / 128;
const int n_col_blocks = (K16 + 3) / 4;
const int out_size_a = n_row_blocks_a * n_col_blocks * 512;
const int out_size_b = n_row_blocks_b * n_col_blocks * 512;
auto scale_a = torch::empty({out_size_a}, sfa.options());
auto scale_b = torch::empty({out_size_b}, sfb.options());
const int total = std::max(M * K16, N * K16);
const int threads = 256;
const int blocks = (total + threads - 1) / threads;
process_scales_kernel<<<blocks, threads>>>(
sfa.data_ptr<uint8_t>(),
sfb.data_ptr<uint8_t>(),
scale_a.data_ptr<uint8_t>(),
scale_b.data_ptr<uint8_t>(),
M, N, K16,
n_row_blocks_a, n_col_blocks, n_row_blocks_b
);
return {scale_a, scale_b};
}
"""
cpp_source = """
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> process_scales(torch::Tensor sfa, torch::Tensor sfb);
"""
# Compile CUDA module
cuda_module = load_inline(
name='grouped_gemm_scales_opt',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['process_scales'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math', '-lineinfo']
)
def custom_kernel(data):
"""Main entry point for the grouped blockscaled GEMM kernel."""
if not isinstance(data, (list, tuple)):
raise ValueError(f"Expected tuple/list, got {type(data)}")
n_items = len(data)
if n_items == 3:
first, second, third = data
if isinstance(first, (list, tuple)) and isinstance(second, (list, tuple)):
return _process_grouped_format(data[0], data[1], data[2])
if n_items == 4:
if all(isinstance(x, (list, tuple)) for x in data):
list0, list1, list2, list3 = data
if len(list0) > 0:
t0 = list0[0]
t1 = list1[0] if len(list1) > 0 else None
t3 = list3[0] if len(list3) > 0 else None
len0 = len(t0) if isinstance(t0, (list, tuple)) else -1
len1 = len(t1) if isinstance(t1, (list, tuple)) else -1
len3 = len(t3) if isinstance(t3, (list, tuple)) else -1
if len0 == 3 and len1 == 2 and len3 == 4:
return _process_grouped_format(list0, list1, list3)
elif len0 == 3 and len1 == 2:
t2 = list2[0] if len(list2) > 0 else None
len2 = len(t2) if isinstance(t2, (list, tuple)) else -1
if len2 == 4:
return _process_grouped_format(list0, list1, list2)
raise ValueError(f"Unrecognized data format with {n_items} elements")
def _process_grouped_format(abc_tensors, sfasfb_tensors, problem_sizes):
"""Process data in grouped format with optimized torch._scaled_mm calls."""
num_groups = len(problem_sizes)
device = abc_tensors[0][0].device
fp8_dtype = sfasfb_tensors[0][0].dtype
# Pre-extract all data to minimize Python overhead in the hot loop
group_data = []
for i in range(num_groups):
a, b, c = abc_tensors[i]
sfa, sfb = sfasfb_tensors[i]
M, N, K, L = problem_sizes[i]
if L == 1:
# Pre-slice and prepare all tensors
a_slice = a[:, :, 0]
b_slice = b[:, :, 0]
sfa_slice = sfa[:, :, 0]
sfb_slice = sfb[:, :, 0]
c_slice = c[:, :, 0]
# Ensure contiguous
if not a_slice.is_contiguous():
a_slice = a_slice.contiguous()
if not b_slice.is_contiguous():
b_slice = b_slice.contiguous()
# Prepare scale factor bytes
sfa_bytes = sfa_slice.contiguous().view(torch.uint8)
sfb_bytes = sfb_slice.contiguous().view(torch.uint8)
# Move to device if needed
if sfa_bytes.device != device:
sfa_bytes = sfa_bytes.to(device)
if sfb_bytes.device != device:
sfb_bytes = sfb_bytes.to(device)
# Pre-transpose B
b_t = b_slice.t()
group_data.append((a_slice, b_t, sfa_bytes, sfb_bytes, c_slice, 1))
else:
group_data.append((a, b, sfa, sfb, c, L))
# Process all groups - hot loop with minimal Python overhead
for i in range(num_groups):
data = group_data[i]
L = data[5]
if L == 1:
a_slice, b_t, sfa_bytes, sfb_bytes, c_slice, _ = data
# Convert scales to blocked format and execute GEMM
scales = cuda_module.process_scales(sfa_bytes, sfb_bytes)
scale_a = scales[0].view(fp8_dtype)
scale_b = scales[1].view(fp8_dtype)
# GEMM: C = A @ B^T with block scaling
result = torch._scaled_mm(a_slice, b_t, scale_a, scale_b, out_dtype=torch.float16)
c_slice.copy_(result)
else:
a, b, sfa, sfb, c, L = data
_process_multi_layer(a, b, sfa, sfb, c, device, L, fp8_dtype)
return [abc_tensors[i][2] for i in range(num_groups)]
def _process_multi_layer(a, b, sfa, sfb, c, device, L, fp8_dtype):
"""Process multi-layer (L > 1) case."""
for l_idx in range(L):
a_slice = a[:, :, l_idx]
b_slice = b[:, :, l_idx]
sfa_slice = sfa[:, :, l_idx]
sfb_slice = sfb[:, :, l_idx]
if not a_slice.is_contiguous():
a_slice = a_slice.contiguous()
if not b_slice.is_contiguous():
b_slice = b_slice.contiguous()
sfa_bytes = sfa_slice.contiguous().view(torch.uint8)
sfb_bytes = sfb_slice.contiguous().view(torch.uint8)
if sfa_bytes.device != device:
sfa_bytes = sfa_bytes.to(device)
if sfb_bytes.device != device:
sfb_bytes = sfb_bytes.to(device)
scales = cuda_module.process_scales(sfa_bytes, sfb_bytes)
scale_a = scales[0].view(fp8_dtype)
scale_b = scales[1].view(fp8_dtype)
b_t = b_slice.t()
result = torch._scaled_mm(a_slice, b_t, scale_a, scale_b, out_dtype=torch.float16)
c[:, :, l_idx].copy_(result)
scrolls · 294 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