Skip to content
KernelIndex
Search⌘K

submission 256645

Batuhanaktas · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

kai_optimizer.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-256645?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
30.5µs
#251 of 420
2026-01-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f6ded74b8aaa532f99520e59bd6af2ba691f7affd52596bde71c5abc7a4d7da8
license declaredunknown
license concludedunknown
authorsBatuhanaktas
imported2026-08-26

Techniques

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

fused-epiloguedef _fused_dual_epilogue_kernel(
num-warps = 8num_warps=8,

Kernel source

kai_optimizer.py86 lines
import torch
import triton
import triton.language as tl
from typing import Tuple

input_t = Tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
]
output_t = torch.Tensor


@triton.jit
def _fused_dual_epilogue_kernel(
    g1_ptr, g2_ptr, out_ptr,
    M, N, L,
    stride_gm, stride_gn,
    stride_om, stride_on, stride_ol,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    total_elements = M * N * L
    offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    mask = offsets < total_elements

    idx_l = offsets % L
    idx_n = (offsets // L) % N
    idx_m = offsets // (L * N)

    gemm_row = idx_l * M + idx_m
    gemm_off = gemm_row * stride_gm + idx_n * stride_gn

    v1 = tl.load(g1_ptr + gemm_off, mask=mask).to(tl.float32)
    v2 = tl.load(g2_ptr + gemm_off, mask=mask).to(tl.float32)
    res = (v1 / (1.0 + tl.exp(-v1))) * v2

    out_off = idx_m * stride_om + idx_n * stride_on + idx_l * stride_ol
    tl.store(out_ptr + out_off, res.to(tl.float16), mask=mask)


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_p, sfb1_p, sfb2_p, c_out = data
    M, K, L = a.shape
    N = b1.shape[0]

    sa = sfa_p.permute(5, 2, 4, 0, 1, 3).reshape(-1)
    sb1 = sfb1_p.permute(5, 2, 4, 0, 1, 3).reshape(-1)
    sb2 = sfb2_p.permute(5, 2, 4, 0, 1, 3).reshape(-1)

    a_f = a.permute(2, 0, 1).reshape(-1, K)
    b1_f = b1.permute(2, 0, 1).reshape(-1, K)
    b2_f = b2.permute(2, 0, 1).reshape(-1, K)

    g1 = torch._scaled_mm(a_f, b1_f.t(), sa, sb1, out_dtype=torch.float32)
    g2 = torch._scaled_mm(a_f, b2_f.t(), sa, sb2, out_dtype=torch.float32)

    total_elements = M * N * L
    BLOCK_SIZE = 2048
    grid = (triton.cdiv(total_elements, BLOCK_SIZE),)

    _fused_dual_epilogue_kernel[grid](
        g1, g2, c_out,
        M, N, L,
        g1.stride(0), g1.stride(1),
        c_out.stride(0), c_out.stride(1), c_out.stride(2),
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=8,
    )

    return c_out


def run_kernel(data: input_t, *, sync: bool = True) -> torch.Tensor:
    result = custom_kernel(data)
    if sync and result.is_cuda:
        torch.cuda.synchronize(result.device)
    return result.contiguous()
scrolls · 86 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