Skip to content
KernelIndex
Search⌘K

submission 148348

Arseni Ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

triton_hardcoded_not_persistent.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-148348?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.2µs
#162 of 369
2025-12-12

Reported · How evidence levels are derived →

Source and license

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

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 = 512BLOCK_K = 512
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128
warp-specializationWARP_SPECIALIZE_INNER: tl.constexpr,

Kernel source

triton_hardcoded_not_persistent.py205 lines
#!POPCORN leaderboard nvfp4_gemm
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor

@triton.jit
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_INNER_STAGES: tl.constexpr,
    WARP_SPECIALIZE_INNER: 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_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid // num_pid_n
    pid_n = pid % 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 constants
    BLOCK_M = 128
    BLOCK_N = 128
    BLOCK_K = 512
    GROUP_SZ = 16
    ELEM_PER_BYTE = 2
    REP_K = BLOCK_K // GROUP_SZ // 4

    # --- Manual Configuration Selection ---
    # Default config (fallback)
    num_inner_stages = 2
    warp_specialize_inner = False
    num_stages = 3
    num_warps = 4

    # Match specific shapes
    if M == 128 and N == 7168 and K == 16384:
        # Config [128x7168x16384]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:2 Wrp:4
        num_inner_stages = 3
        warp_specialize_inner = True
        num_stages = 2
        num_warps = 4
    elif M == 128 and N == 4096 and K == 7168:
        # Config [128x4096x7168]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:3 Wrp:4
        num_inner_stages = 3
        warp_specialize_inner = True
        num_stages = 2
        num_warps = 4
    elif M == 128 and N == 7168 and K == 2048:
        # Config [128x7168x2048]: BK:512 OS:3 IS:2 WSO:False WSI:False Stg:2 Wrp:4
        num_inner_stages = 3
        warp_specialize_inner = False
        num_stages = 2
        num_warps = 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)
    grid=(num_tiles,)
    
    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,
        NUM_INNER_STAGES=num_inner_stages,
        WARP_SPECIALIZE_INNER=warp_specialize_inner,
        num_warps=num_warps,
        num_stages=num_stages
    )

    return c_tensor
scrolls · 205 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 145523.

⋯ 21 unchanged lines
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
+ pid_m = pid // num_pid_n
+ pid_n = pid % 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 linear in tl.range(
- pid,
- total_tiles,
- num_pid,
- num_stages=NUM_OUTER_STAGES,
- flatten=FLATTEN,
- warp_specialize=WARP_SPECIALIZE_OUTER,
+ for i in tl.range(
+ 0,
+ tl.cdiv(K, BLOCK_K),
+ num_stages=NUM_INNER_STAGES,
+ warp_specialize=WARP_SPECIALIZE_INNER,
):
- tile_id = linear % (num_pid_m * num_pid_n)
- pid_m = tile_id // num_pid_n
- pid_n = tile_id % num_pid_n
+ offs_k = i * BLOCK_K_ELEM_PER_BYTE
+ offs_scale_k = i * REP_K
- # Base offsets for this tile
- offs_am = pid_m * BLOCK_M
- offs_bn = pid_n * BLOCK_N
+ # 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])
- accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ 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)
+ )
- 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
+ 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)
+ )
- # 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
+ accumulator = tl.dot_scaled(
+ a,
+ scale_a,
+ "e2m1",
+ b.T,
+ scale_b,
+ "e2m1",
+ accumulator,
)
-
- c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)
- tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)
+ # 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
⋯ 18 unchanged lines
# --- Manual Configuration Selection ---
# Default config (fallback)
- num_outer_stages = 2
num_inner_stages = 2
- warp_specialize_outer = True
warp_specialize_inner = False
num_stages = 3
num_warps = 4
⋯ 1 unchanged lines
# Match specific shapes
if M == 128 and N == 7168 and K == 16384:
# Config [128x7168x16384]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:2 Wrp:4
- num_outer_stages = 2
num_inner_stages = 3
- warp_specialize_outer = True
warp_specialize_inner = True
num_stages = 2
num_warps = 4
elif M == 128 and N == 4096 and K == 7168:
# Config [128x4096x7168]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:3 Wrp:4
- num_outer_stages = 2
num_inner_stages = 3
- warp_specialize_outer = True
warp_specialize_inner = True
- num_stages = 3
+ num_stages = 2
num_warps = 4
elif M == 128 and N == 7168 and K == 2048:
# Config [128x7168x2048]: BK:512 OS:3 IS:2 WSO:False WSI:False Stg:2 Wrp:4
- num_outer_stages = 3
- num_inner_stages = 2
- warp_specialize_outer = False
+ num_inner_stages = 3
warp_specialize_inner = False
num_stages = 2
num_warps = 4
⋯ 56 unchanged lines
BLOCK_N,
BLOCK_K,
REP_K,
- NUM_OUTER_STAGES=num_outer_stages,
NUM_INNER_STAGES=num_inner_stages,
- WARP_SPECIALIZE_OUTER=warp_specialize_outer,
WARP_SPECIALIZE_INNER=warp_specialize_inner,
- FLATTEN=True,
num_warps=num_warps,
num_stages=num_stages
)
scrolls · 191 diff lines total

Best evidence level for this revision: reported

JSON