Skip to content
KernelIndex
Search⌘K

submission 213657

Arseni Ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

triton_dual_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-213657?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 dual GEMMsuite of 4 cases
NVIDIA B200
21.2µs
#187 of 420
2025-12-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9cfe6c8bde81bc7d39708da671e2eaa93258f7d26623ae076fc0c515cef6f95d
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 = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128
warp-specializationWARP_SPECIALIZE_INNER: tl.constexpr,

Kernel source

triton_dual_gemm.py236 lines
#!POPCORN leaderboard nvfp4_dual_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,
    b1_desc,
    b2_desc,
    b1_scale_desc,
    b2_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_M: tl.constexpr,
    REP_N: tl.constexpr,
    REP_K: tl.constexpr,
    NUM_INNER_STAGES: tl.constexpr,
    WARP_SPECIALIZE_INNER: 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

    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

    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

    accumulator1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)
    accumulator2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)

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

        a = a_desc.load([offs_am, 0, offs_k])
        b1 = b1_desc.load([offs_bn, 0, offs_k])
        b2 = b2_desc.load([offs_bn, 0, offs_k])
        
        a = a.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)
        b1 = b1.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)
        b2 = b2.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)

        scale_a = (
            a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
            .reshape(REP_M, REP_K, 32, 4, 4)
            .trans(0, 3, 2, 1, 4)
            .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
        )
        scale_b1 = (
            b1_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
            .reshape(REP_N, REP_K, 32, 4, 4)
            .trans(0, 3, 2, 1, 4)
            .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
        )
        scale_b2 = (
            b2_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
            .reshape(REP_N, REP_K, 32, 4, 4)
            .trans(0, 3, 2, 1, 4)
            .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
        )
        
        accumulator1 = tl.dot_scaled(
            a,
            scale_a,
            "e2m1",
            b1.T,
            scale_b1,
            "e2m1",
            accumulator1,
        )
        accumulator2 = tl.dot_scaled(
            a,
            scale_a,
            "e2m1",
            b2.T,
            scale_b2,
            "e2m1",
            accumulator2,
        )

    temp1 = accumulator1 * (1 / (1 + tl.exp(-accumulator1)))
    accumulator2 = temp1 * accumulator2

    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, accumulator2.to(output_dtype), mask=c_mask, cache_modifier=".cg")


def custom_kernel(data):
    a_tensor, b1_tensor, b2_tensor, _, _, _, sfa_tensor, sfb1_tensor, sfb2_tensor, c = data
    
    M, K_half, L = a_tensor.shape
    N = b1_tensor.shape[0]
    K = K_half * 2
    
    BLOCK_M = 128
    BLOCK_N = 128
    BLOCK_K = 256
    num_stages = 3
    num_warps = 4
    warp_specialize_inner = False
    
    if M >= 256 and N >= 256:
        BLOCK_M = 128
        BLOCK_N = 128
    
    if K >= 4096:
        BLOCK_K = 256
        num_stages = 3
        warp_specialize_inner = True
        num_warps = 4
    elif K >= 2048:
        BLOCK_K = 128
        num_stages = 4
    else:
        BLOCK_K = 128
        
    if M == 256 and N == 4096:
        BLOCK_M = 128; BLOCK_N = 128; BLOCK_K = 256
        num_stages = 3; num_warps = 4; warp_specialize_inner = True

    GROUP_SZ = 16
    ELEM_PER_BYTE = 2
    
    REP_M = BLOCK_M // 128
    REP_N = BLOCK_N // 128
    REP_K = BLOCK_K // GROUP_SZ // 4

    a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1) 
    a_desc = TensorDescriptor.from_tensor(
        a_tma,
        block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],
    )
    
    b1_tma = b1_tensor.view(torch.uint8).permute(0, 2, 1) 
    b2_tma = b2_tensor.view(torch.uint8).permute(0, 2, 1)
    b1_desc = TensorDescriptor.from_tensor(
        b1_tma,
        block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],
    )
    b2_desc = TensorDescriptor.from_tensor(
        b2_tma,
        block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],
    )

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

    sfa_back = sfa_tensor.permute(5, 2, 4, 0, 1, 3)
    sfb1_back = sfb1_tensor.permute(5, 2, 4, 0, 1, 3)
    sfb2_back = sfb2_tensor.permute(5, 2, 4, 0, 1, 3)

    a_scale_packed = sfa_back.view(L, rest_m, rest_k, 2, 256)
    b1_scale_packed = sfb1_back.view(L, rest_n, rest_k, 2, 256)
    b2_scale_packed = sfb2_back.view(L, rest_n, rest_k, 2, 256)
    
    a_scale_desc = TensorDescriptor.from_tensor(
        a_scale_packed,
        block_shape=[1, REP_M, REP_K, 2, 256],
    )
    b1_scale_desc = TensorDescriptor.from_tensor(
        b1_scale_packed,
        block_shape=[1, REP_N, REP_K, 2, 256],
    )
    b2_scale_desc = TensorDescriptor.from_tensor(
        b2_scale_packed,
        block_shape=[1, REP_N, REP_K, 2, 256],
    )

    stride_cm, stride_cn, _ = c.stride()

    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,
        b1_desc,
        b2_desc,
        b1_scale_desc,
        b2_scale_desc,
        c,
        stride_cm,
        stride_cn,
        M,
        N,
        K,
        ELEM_PER_BYTE,
        GROUP_SZ,
        BLOCK_M,
        BLOCK_N,
        BLOCK_K,
        REP_M,
        REP_N,
        REP_K,
        NUM_INNER_STAGES=num_stages,
        WARP_SPECIALIZE_INNER=warp_specialize_inner,
        num_warps=num_warps,
        num_stages=num_stages
    )

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