Skip to content
KernelIndex
Search⌘K

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
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
804.7µs
#133 of 145
2026-01-21

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.

fp4This implementation computes grouped matrix multiplication with FP4 inputs

Kernel 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