Skip to content
KernelIndex
Search⌘K

submission 126350

nrehiew · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 124 lines, June 9 Researcher Reciprocity License v1.0.

submission_naive.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-126350?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 GEMMsuite of 3 cases
NVIDIA B200
20.7µs
#194 of 369
2025-12-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ed18112f91f0171c5cd41c3cffde3dada0ee907b6f9e6075a004e0b451ac9c30
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 4num_warps=4,
stages = 3num_stages=3,
tile-k = 256BLOCK_SIZE_K = 256
tile-m = 128BLOCK_SIZE_M = 128
tile-n = 128BLOCK_SIZE_N = 128

Kernel source

submission_naive.py124 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
from task import input_t, output_t

sf_vec_size = 16
elements_per_byte = 2

ACC_DTYPE = tl.float32


@triton.jit
def triton_kernel(
    a_desc,
    b_desc,
    sfa_desc,
    sfb_desc,
    c_ptr,
    c_m_stride,
    c_n_stride,
    c_l_stride,
    M,
    K,
    BLOCK_SIZE_K: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    rep_m: tl.constexpr,
    rep_n: tl.constexpr,
    rep_k: tl.constexpr,
    elements_per_byte: tl.constexpr = elements_per_byte,
    sf_vec_size: tl.constexpr = sf_vec_size,
):
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    pid_m = pid % num_pid_m
    pid_n = pid // num_pid_m
    offs_am = pid_m * BLOCK_SIZE_M
    offs_bn = pid_n * BLOCK_SIZE_N
    offs_k_a = 0
    offs_k_b = 0
    offs_scale_m = pid_m * rep_m
    offs_scale_n = pid_n * rep_n
    offs_scale_k = 0

    packed_k = (K + elements_per_byte - 1) // elements_per_byte
    acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=ACC_DTYPE)

    for k_byte in tl.range(0, packed_k, BLOCK_SIZE_K // elements_per_byte, num_stages=num_stages, disallow_acc_multi_buffer=True, flatten=True):
        a_val_uint8 = a_desc.load([0, offs_am, offs_k_a]).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte)
        b_val_uint8 = b_desc.load([0, offs_bn, offs_k_b]).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte)
        a_scale_val = sfa_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
        b_scale_val = sfb_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])

        a_scale_val = a_scale_val.reshape(rep_m, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // sf_vec_size)
        b_scale_val = b_scale_val.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // sf_vec_size)
        acc = tl.dot_scaled(a_val_uint8, a_scale_val, "e2m1", b_val_uint8.T, b_scale_val, "e2m1", acc)

        offs_k_a += BLOCK_SIZE_K // elements_per_byte
        offs_k_b += BLOCK_SIZE_K // elements_per_byte
        offs_scale_k += rep_k

    out_offsets = (offs_am + tl.arange(0, BLOCK_SIZE_M))[:, None] * c_m_stride + (offs_bn + tl.arange(0, BLOCK_SIZE_N))[None, :] * c_n_stride
    tl.store(c_ptr + out_offsets, acc.to(tl.float16))


def custom_kernel(data: input_t) -> output_t:
    # c: [m, n, l] is pre-allocated memory to avoid timing allocation overhead.
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
    m, k_packed, l = a.shape
    k = k_packed * elements_per_byte
    m, n, l = c.shape
    # a: [128, K, 1]
    # b: [N, K, 1]
    # sfa: [128, K // 16, 1]
    # sfb: [N, K // 16, 1]
    # c: [M, N, 1]
    BLOCK_SIZE_M = 128
    BLOCK_SIZE_K = 256
    BLOCK_SIZE_N = 128

    rep_m = BLOCK_SIZE_M // 128
    rep_n = BLOCK_SIZE_N // 128
    # rep_n = n // 128
    rep_k = BLOCK_SIZE_K // sf_vec_size // 4

    a = a.view(torch.uint8).permute(2, 0, 1)  # [l, m, k]
    b = b.view(torch.uint8).permute(2, 0, 1)  # [l, n, k]
    sfa_permuted_permute = sfa_permuted.permute(5, 2, 4, 0, 1, 3)
    sfb_permuted_permute = sfb_permuted.permute(5, 2, 4, 0, 1, 3)
    sfa_5d = sfa_permuted_permute.reshape(l, sfa_permuted_permute.shape[1], sfa_permuted_permute.shape[2], 2, 256)
    sfb_5d = sfb_permuted_permute.reshape(l, sfb_permuted_permute.shape[1], sfb_permuted_permute.shape[2], 2, 256)

    a_desc = TensorDescriptor.from_tensor(a, [1, BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte])
    b_desc = TensorDescriptor.from_tensor(b, [1, BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte])
    a_scale_desc = TensorDescriptor.from_tensor(sfa_5d, [1, rep_m, rep_k, 2, 256])
    b_scale_desc = TensorDescriptor.from_tensor(sfb_5d, [1, rep_n, rep_k, 2, 256])

    grid = lambda meta: (triton.cdiv(m, meta["BLOCK_SIZE_M"]) * triton.cdiv(n, meta["BLOCK_SIZE_N"]), l)
    triton_kernel[grid](
        a_desc,
        b_desc,
        a_scale_desc,
        b_scale_desc,
        c,
        c.stride(0),
        c.stride(1),
        c.stride(2),
        m,
        k,
        BLOCK_SIZE_K=BLOCK_SIZE_K,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        num_warps=4,
        num_stages=3,
        rep_m=rep_m,
        rep_n=rep_n,
        rep_k=rep_k,
    )

    return c
scrolls · 124 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