Skip to content
KernelIndex
Search⌘K

submission 303620

VladRad · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-303620?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
24.4µs
#207 of 420
2026-01-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0a6d63322174ccf6672b1f0f17fd4a234325a764f190a802eea0793d9e5d04bf
license declaredunknown
license concludedunknown
authorsVladRad
imported2026-08-26

Techniques

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

fp4"""NVFP4 Dual GEMM + SiLU - Working Optimized Triton
num-warps = 4num_warps=4,
stages = 4num_stages=4,
tile-k = 256BLOCK_SIZE_K = 256
tile-m = 128BLOCK_SIZE_M = 128
tile-n = 128BLOCK_SIZE_N = 128

Kernel source

submission.py140 lines
"""NVFP4 Dual GEMM + SiLU - Working Optimized Triton
C = silu(A @ B1) * (A @ B2)

This version passes all tests and provides good performance.
"""

from __future__ import annotations

import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor


@triton.jit
def dual_gemm_silu_kernel(
    a_desc, b1_desc, b2_desc,
    sfa_desc, sfb1_desc, sfb2_desc,
    c_ptr, c_m_stride, c_n_stride, c_l_stride,
    M, K,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    rep_m: tl.constexpr,
    rep_n: tl.constexpr,
    rep_k: tl.constexpr,
    num_stages: tl.constexpr,
    sf_vec_size: tl.constexpr,
    elements_per_byte: tl.constexpr,
):
    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 = 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

    acc1 = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    acc2 = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    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_tile = a_desc.load([0, offs_am, offs_k]).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte)
        b1_tile = b1_desc.load([0, offs_bn, offs_k]).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte)
        b2_tile = b2_desc.load([0, offs_bn, offs_k]).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte)

        sfa_tile = sfa_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
        sfb1_tile = sfb1_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
        sfb2_tile = sfb2_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])

        sfa_tile = sfa_tile.reshape(rep_m, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // sf_vec_size)
        sfb1_tile = sfb1_tile.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // sf_vec_size)
        sfb2_tile = sfb2_tile.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // sf_vec_size)

        acc1 = tl.dot_scaled(a_tile, sfa_tile, "e2m1", b1_tile.T, sfb1_tile, "e2m1", acc1)
        acc2 = tl.dot_scaled(a_tile, sfa_tile, "e2m1", b2_tile.T, sfb2_tile, "e2m1", acc2)

        offs_k += BLOCK_SIZE_K // elements_per_byte
        offs_scale_k += rep_k

    LOG2E = tl.full((1,), 1.4426950408889634, dtype=tl.float32)
    exp_neg = tl.exp2(-acc1 * LOG2E)
    sigmoid1 = 1.0 / (1.0 + exp_neg)
    result = acc1 * sigmoid1 * acc2

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


def custom_kernel(data: tuple) -> torch.Tensor:
    a, b1, b2, sfa, sfb1, sfb2, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
    
    elements_per_byte = 2
    sf_vec_size = 16
    
    m, k_packed, l = a.shape
    k = k_packed * elements_per_byte
    n = c.shape[1]
    
    # Fixed block sizes
    BLOCK_SIZE_M = 128
    BLOCK_SIZE_N = 128
    BLOCK_SIZE_K = 256
    
    rep_m = BLOCK_SIZE_M // 128  # 1
    rep_n = BLOCK_SIZE_N // 128  # 1
    rep_k = BLOCK_SIZE_K // sf_vec_size // 4  # 4
    
    # Data prep
    a_tma = a.view(torch.uint8).permute(2, 0, 1).contiguous()
    b1_tma = b1.view(torch.uint8).permute(2, 0, 1).contiguous()
    b2_tma = b2.view(torch.uint8).permute(2, 0, 1).contiguous()
    
    sfa_5d = sfa_permuted.permute(5, 2, 4, 0, 1, 3).reshape(
        l, sfa_permuted.shape[2], sfa_permuted.shape[4], 2, 256
    ).contiguous()
    sfb1_5d = sfb1_permuted.permute(5, 2, 4, 0, 1, 3).reshape(
        l, sfb1_permuted.shape[2], sfb1_permuted.shape[4], 2, 256
    ).contiguous()
    sfb2_5d = sfb2_permuted.permute(5, 2, 4, 0, 1, 3).reshape(
        l, sfb2_permuted.shape[2], sfb2_permuted.shape[4], 2, 256
    ).contiguous()
    
    # Descriptors
    a_desc = TensorDescriptor.from_tensor(a_tma, [1, BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte])
    b1_desc = TensorDescriptor.from_tensor(b1_tma, [1, BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte])
    b2_desc = TensorDescriptor.from_tensor(b2_tma, [1, BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte])
    sfa_desc = TensorDescriptor.from_tensor(sfa_5d, [1, rep_m, rep_k, 2, 256])
    sfb1_desc = TensorDescriptor.from_tensor(sfb1_5d, [1, rep_n, rep_k, 2, 256])
    sfb2_desc = TensorDescriptor.from_tensor(sfb2_5d, [1, rep_n, rep_k, 2, 256])
    
    grid = (triton.cdiv(m, BLOCK_SIZE_M) * triton.cdiv(n, BLOCK_SIZE_N),)
    dual_gemm_silu_kernel[grid](
        a_desc, b1_desc, b2_desc,
        sfa_desc, sfb1_desc, sfb2_desc,
        c, c.stride(0), c.stride(1), c.stride(2),
        m, k,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        BLOCK_SIZE_K=BLOCK_SIZE_K,
        rep_m=rep_m,
        rep_n=rep_n,
        rep_k=rep_k,
        num_stages=4,
        sf_vec_size=sf_vec_size,
        elements_per_byte=elements_per_byte,
        num_warps=4,
    )
    
    return c
scrolls · 140 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