Skip to content
KernelIndex
Search⌘K

submission 144525

Arseni Ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

triton_naive.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-144525?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
16.6µs
#165 of 369
2025-12-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c6fe0f6fe035374a7f580b9b6c55b3cd4f1680e253142fefab42248354676314
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-15

Techniques

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

autotunedef _config(**autotune_kwargs):
num-warps = 4num_warps=4,
persistent-kernelnum_pid = tl.num_programs(axis=0)
stages = 3num_stages=3,
tile-k = 512BLOCK_K = 512
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128
warp-specializationWARP_SPECIALIZE_OUTER=True,

Kernel source

triton_naive.py225 lines
#!POPCORN leaderboard nvfp4_gemm
import functools
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor

def _matmul_launch_metadata(grid, kernel, args):
    M, N, K = args["M"], args["N"], args["K"]
    return {
        "name": f"{kernel.name} [M={M}, N={N}, K={K}]",
        "flops": 2.0 * M * N * K,
    }


def _config(**autotune_kwargs):
    class inner:
        def __init__(self, fn):
            self.fn = fn

        def __getitem__(self, s):
            return functools.partial(self.fn[s], **autotune_kwargs)

    return inner


@_config(
    NUM_OUTER_STAGES=None,
    NUM_INNER_STAGES=None,
    WARP_SPECIALIZE_OUTER=True,
    WARP_SPECIALIZE_INNER=False,
    FLATTEN=True,
    num_warps=4,
    num_stages=3,
    num_ctas=1,
)
@triton.jit(launch_metadata=_matmul_launch_metadata)
def block_scaled_batched_gemm_kernel(
    a_desc,
    a_scale_desc,
    b_desc,
    b_scale_desc,
    c_ptr,
    stride_cm,
    stride_cn,
    M,
    N,
    K,
    ELEM_PER_BYTE: tl.constexpr,
    GROUP_SZ: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    REP_K: tl.constexpr,
    NUM_OUTER_STAGES: tl.constexpr,
    NUM_INNER_STAGES: tl.constexpr,
    WARP_SPECIALIZE_OUTER: tl.constexpr,
    WARP_SPECIALIZE_INNER: tl.constexpr,
    FLATTEN: tl.constexpr,
):
    BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTE
    BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ

    pid = tl.program_id(axis=0)
    num_pid = tl.num_programs(axis=0)

    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    total_tiles = num_pid_m * num_pid_n

    for linear in tl.range(
        pid,
        total_tiles,
        num_pid,
        num_stages=NUM_OUTER_STAGES,
        flatten=FLATTEN,
        warp_specialize=WARP_SPECIALIZE_OUTER,
    ):
        tile_id = linear % (num_pid_m * num_pid_n)
        pid_m = tile_id // num_pid_n
        pid_n = tile_id % num_pid_n

        # Base offsets for this tile
        offs_am = pid_m * BLOCK_M
        offs_bn = pid_n * BLOCK_N

        accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

        for i in tl.range(
            0,
            tl.cdiv(K, BLOCK_K),
            num_stages=NUM_INNER_STAGES,
            warp_specialize=WARP_SPECIALIZE_INNER,
        ):
            offs_k = i * BLOCK_K_ELEM_PER_BYTE
            offs_scale_k = i * REP_K

            # A: [BLOCK_M, BLOCK_K/2]
            # B: [BLOCK_N, BLOCK_K/2]
            a = a_desc.load([offs_am, offs_k])
            b = b_desc.load([offs_bn, offs_k])

            scale_a = (
                a_scale_desc.load([pid_m, offs_scale_k, 0, 0])
                .reshape(REP_K, 32, 4, 4)
                .trans(2, 1, 0, 3)
                .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
            )

            scale_b = (
                b_scale_desc.load([pid_n, offs_scale_k, 0, 0])
                .reshape(REP_K, 32, 4, 4)
                .trans(2, 1, 0, 3)
                .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
            )
            accumulator = tl.dot_scaled(
                a,
                scale_a,
                "e2m1",
                b.T,
                scale_b,
                "e2m1",
                accumulator,
            )

        # Calculate output pointers
        offset_m = offs_am + tl.arange(0, BLOCK_M)
        offset_n = offs_bn + tl.arange(0, BLOCK_N)
        
        c_off = (
            offset_m[:, None] * stride_cm
            + offset_n[None, :] * stride_cn
        )
        
        c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)
        tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)


def custom_kernel(data):
    a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data
    
    #We only have a single batch every time
    a_tensor = a_tensor.squeeze(-1)
    b_tensor = b_tensor.squeeze(-1)
    sfa_tensor = sfa_tensor.squeeze(-1)
    sfb_tensor = sfb_tensor.squeeze(-1)

    # Input Shapes
    M, K_half = a_tensor.shape
    N = b_tensor.shape[0]
    K = K_half * 2
    
    # Configuration
    BLOCK_M = 128
    BLOCK_N = 128
    BLOCK_K = 512
    GROUP_SZ = 16
    ELEM_PER_BYTE = 2
    SM_MULT = 1
    REP_K = BLOCK_K // GROUP_SZ // 4

    a_tma = a_tensor.view(torch.uint8) # [M, K/2]
    a_desc = TensorDescriptor.from_tensor(
        a_tma,
        block_shape=[BLOCK_M, BLOCK_K // ELEM_PER_BYTE],
    )
    
    b_tma = b_tensor.view(torch.uint8) # [N, K/2]
    b_desc = TensorDescriptor.from_tensor(
        b_tma,
        block_shape=[BLOCK_N, BLOCK_K // ELEM_PER_BYTE],
    )

    rest_m = M // 128
    rest_n = N // 128
    rest_k = triton.cdiv(K, GROUP_SZ) // 4

    # sfa_permuted: [32, 4, rest_m, 4, rest_k]
    # sfb_permuted: [32, 4, rest_n, 4, rest_k]
    # Permute to [rest_m or rest_n, rest_k, 32, 4, 4]
    sfa_back = sfa_tensor.permute(2, 4, 0, 1, 3)
    sfb_back = sfb_tensor.permute(2, 4, 0, 1, 3)

    # Pack final three dims: (rest_m, rest_k, 32, 4, 4) -> (rest_m, rest_k, 2, 256)
    a_scale_packed = sfa_back.view(rest_m, rest_k, 2, 256)
    b_scale_packed = sfb_back.view(rest_n, rest_k, 2, 256)
    a_scale_desc = TensorDescriptor.from_tensor(
        a_scale_packed,
        block_shape=[1, REP_K, 2, 256],
    )
    
    b_scale_desc = TensorDescriptor.from_tensor(
        b_scale_packed,
        block_shape=[1, REP_K, 2, 256],
    )

    stride_cm, stride_cn, _ = c_tensor.stride()
    
    # Launch Grid
    num_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)
    num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count
    # Persistent kernel grid size
    grid = (min(num_tiles, num_sms * SM_MULT),)

    block_scaled_batched_gemm_kernel[grid](
        a_desc,
        a_scale_desc,
        b_desc,
        b_scale_desc,
        c_tensor,
        stride_cm,
        stride_cn,
        M,
        N,
        K,
        ELEM_PER_BYTE,
        GROUP_SZ,
        BLOCK_M,
        BLOCK_N,
        BLOCK_K,
        REP_K,
    )

    return c_tensor
scrolls · 225 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 120588.

⋯ 42 unchanged lines
c_ptr,
stride_cm,
stride_cn,
- stride_cl,
M,
N,
K,
⋯ 2 unchanged lines
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
- REP_M: tl.constexpr,
- REP_N: tl.constexpr,
REP_K: tl.constexpr,
NUM_OUTER_STAGES: tl.constexpr,
NUM_INNER_STAGES: tl.constexpr,
⋯ 1 unchanged lines
WARP_SPECIALIZE_INNER: tl.constexpr,
FLATTEN: tl.constexpr,
):
- output_dtype: tl.constexpr = tl.float16
- acc_dtype: tl.constexpr = tl.float32
BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTE
BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ
⋯ 12 unchanged lines
flatten=FLATTEN,
warp_specialize=WARP_SPECIALIZE_OUTER,
):
- # Decode linear index into (pid_m, pid_n, pid_b)
tile_id = linear % (num_pid_m * num_pid_n)
- pid_b = linear // (num_pid_m * num_pid_n)
-
pid_m = tile_id // num_pid_n
pid_n = tile_id % num_pid_n
# Base offsets for this tile
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
- offs_scale_m = pid_m * REP_M
- offs_scale_n = pid_n * REP_N
- accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for i in tl.range(
0,
⋯ 9 unchanged lines
a = a_desc.load([offs_am, offs_k])
b = b_desc.load([offs_bn, offs_k])
- # Reconstruct packed scales for A
scale_a = (
- a_scale_desc.load([pid_b, offs_scale_m, offs_scale_k, 0, 0])
- .reshape(REP_M, REP_K, 32, 4, 4)
- .trans(0, 3, 2, 1, 4)
+ a_scale_desc.load([pid_m, offs_scale_k, 0, 0])
+ .reshape(REP_K, 32, 4, 4)
+ .trans(2, 1, 0, 3)
.reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
)
- # Reconstruct packed scales for B (using pid_n/REP_N)
scale_b = (
- b_scale_desc.load([pid_b, offs_scale_n, offs_scale_k, 0, 0])
- .reshape(REP_M, REP_K, 32, 4, 4)
- .trans(0, 3, 2, 1, 4)
+ b_scale_desc.load([pid_n, offs_scale_k, 0, 0])
+ .reshape(REP_K, 32, 4, 4)
+ .trans(2, 1, 0, 3)
.reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
)
-
- # Scaled Dot Product: A * B.T
accumulator = tl.dot_scaled(
a,
scale_a,
⋯ 11 unchanged lines
c_off = (
offset_m[:, None] * stride_cm
+ offset_n[None, :] * stride_cn
- + pid_b * stride_cl
)
c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)
- tl.store(c_ptr + c_off, accumulator.to(output_dtype), mask=c_mask, cache_modifier=".cg")
+ tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)
def custom_kernel(data):
a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data
-
+
+ #We only have a single batch every time
a_tensor = a_tensor.squeeze(-1)
b_tensor = b_tensor.squeeze(-1)
+ sfa_tensor = sfa_tensor.squeeze(-1)
+ sfb_tensor = sfb_tensor.squeeze(-1)
+
# Input Shapes
- # a: [M, K/2, L], b: [N, K/2, L], sfa: [M, K/16, L], sfb: [N, K/16, L]
M, K_half = a_tensor.shape
N = b_tensor.shape[0]
K = K_half * 2
⋯ 1 unchanged lines
# Configuration
BLOCK_M = 128
BLOCK_N = 128
- BLOCK_K = 256
+ BLOCK_K = 512
GROUP_SZ = 16
ELEM_PER_BYTE = 2
SM_MULT = 1
-
- REP_M = BLOCK_M // 128
- REP_N = BLOCK_N // 128
REP_K = BLOCK_K // GROUP_SZ // 4
- # Prepare TMA Descriptors
- # View as uint8 for TMA to handle the 4-bit packed data correctly
a_tma = a_tensor.view(torch.uint8) # [M, K/2]
a_desc = TensorDescriptor.from_tensor(
a_tma,
block_shape=[BLOCK_M, BLOCK_K // ELEM_PER_BYTE],
)
- b_tma = b_tensor.view(torch.uint8)# [N, K/2]
+ b_tma = b_tensor.view(torch.uint8) # [N, K/2]
b_desc = TensorDescriptor.from_tensor(
b_tma,
block_shape=[BLOCK_N, BLOCK_K // ELEM_PER_BYTE],
)
- # Scales: invert CuTe layout and pack for TMA
- # sfa_permuted: [32, 4, rest_m, 4, rest_k, L]
- # sfb_permuted: [32, 4, rest_n, 4, rest_k, L]
rest_m = M // 128
- rest_n = N // 128 # = 1
+ rest_n = N // 128
rest_k = triton.cdiv(K, GROUP_SZ) // 4
- # Permute to [L, rest_m, rest_k, 32, 4, 4]
- # Permute to [rest_m, rest_k, 32, 4, 4]
- sfa_back = sfa_tensor.permute(5, 2, 4, 0, 1, 3)
- sfb_back = sfb_tensor.permute(5, 2, 4, 0, 1, 3)
- assert sfa_back.shape == (1, rest_m, rest_k, 32, 4, 4)
- assert sfb_back.shape == (1, rest_n, rest_k, 32, 4, 4)
+ # sfa_permuted: [32, 4, rest_m, 4, rest_k]
+ # sfb_permuted: [32, 4, rest_n, 4, rest_k]
+ # Permute to [rest_m or rest_n, rest_k, 32, 4, 4]
+ sfa_back = sfa_tensor.permute(2, 4, 0, 1, 3)
+ sfb_back = sfb_tensor.permute(2, 4, 0, 1, 3)
- # Pack final three dims: (L, rest_m, rest_k, 32, 4, 4) -> (L, rest_m, rest_k, 2, 256)
- a_scale_packed = sfa_back.view(1, rest_m, rest_k, 2, 256)
- b_scale_packed = sfb_back.view(1, rest_n, rest_k, 2, 256)
+ # Pack final three dims: (rest_m, rest_k, 32, 4, 4) -> (rest_m, rest_k, 2, 256)
+ a_scale_packed = sfa_back.view(rest_m, rest_k, 2, 256)
+ b_scale_packed = sfb_back.view(rest_n, rest_k, 2, 256)
a_scale_desc = TensorDescriptor.from_tensor(
a_scale_packed,
- block_shape=[1, REP_M, REP_K, 2, 256],
+ block_shape=[1, REP_K, 2, 256],
)
b_scale_desc = TensorDescriptor.from_tensor(
b_scale_packed,
- block_shape=[1, REP_N, REP_K, 2, 256],
+ block_shape=[1, REP_K, 2, 256],
)
- stride_cm, stride_cn, stride_cl = c_tensor.stride()
+ stride_cm, stride_cn, _ = c_tensor.stride()
# Launch Grid
num_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)
⋯ 9 unchanged lines
c_tensor,
stride_cm,
stride_cn,
- stride_cl,
M,
N,
K,
⋯ 2 unchanged lines
BLOCK_M,
BLOCK_N,
BLOCK_K,
- REP_M,
- REP_N,
REP_K,
)
scrolls · 194 diff lines total

Best evidence level for this revision: reported

JSON