Skip to content
KernelIndex
Search⌘K

gpt-5 / triton998d17

gpt-5_triton_998d17 · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 169 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-998d17?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16

Benchmark evidence

43 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n28672 k4096fp16 · [1, 4096]
NVIDIA B200
235.9µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2, 4096]
NVIDIA B200
243.7µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [4, 4096]
NVIDIA B200
244.9µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [7, 4096]
NVIDIA B200
246.1µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8, 4096]
NVIDIA B200
246.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [15, 4096]
NVIDIA B200
248.3µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [16, 4096]
NVIDIA B200
248.3µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [24, 4096]
NVIDIA B200
249.8µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [32, 4096]
NVIDIA B200
251.1µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [35, 4096]
NVIDIA B200
251.7µs
#7 of 8
2025-10-16
Show all 43 measurements ›
GEMM n28672 k4096fp16 · [40, 4096]
NVIDIA B200
252.2µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [48, 4096]
NVIDIA B200
252.8µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [56, 4096]
NVIDIA B200
253.2µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [64, 4096]
NVIDIA B200
253.7µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [70, 4096]
NVIDIA B200
254.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [88, 4096]
NVIDIA B200
254.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [72, 4096]
NVIDIA B200
254.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [80, 4096]
NVIDIA B200
255.1µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [104, 4096]
NVIDIA B200
256.0µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [96, 4096]
NVIDIA B200
256.3µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [120, 4096]
NVIDIA B200
257.0µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [112, 4096]
NVIDIA B200
257.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [128, 4096]
NVIDIA B200
257.8µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [136, 4096]
NVIDIA B200
469.3µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [144, 4096]
NVIDIA B200
471.3µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [152, 4096]
NVIDIA B200
472.8µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [160, 4096]
NVIDIA B200
475.1µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [168, 4096]
NVIDIA B200
476.0µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [176, 4096]
NVIDIA B200
476.6µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [184, 4096]
NVIDIA B200
477.5µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [192, 4096]
NVIDIA B200
478.4µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [200, 4096]
NVIDIA B200
479.0µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [216, 4096]
NVIDIA B200
480.1µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [208, 4096]
NVIDIA B200
480.2µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [240, 4096]
NVIDIA B200
481.1µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [232, 4096]
NVIDIA B200
482.3µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [224, 4096]
NVIDIA B200
482.5µs
#7 of 8
2025-10-16
GEMM n28672 k4096fp16 · [248, 4096]
NVIDIA B200
483.0µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [256, 4096]
NVIDIA B200
484.9µs
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [972, 4096]
NVIDIA B200
1.74ms
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2053, 4096]
NVIDIA B200
2.73ms
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2379, 4096]
NVIDIA B200
3.14ms
#8 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8192, 4096]
NVIDIA B200
11.3ms
#8 of 8
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:3bde1487665689bc3e962fc344c2df1d862c5031ff03314d3d109da0ca27199a
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Techniques

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

autotune@triton.autotune(
mmaacc += tl.dot(a, b)
num-warps = 8triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),
stages = 5triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),

Kernel source

main.py169 lines
import torch
import triton
import triton.language as tl


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=4, num_stages=5),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=4, num_stages=5),
        triton.Config({'BLOCK_M': 256, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 128}, num_warps=8, num_stages=4),
    ],
    key=['M'],
)
@triton.jit
def _gemm_n_28672_k_4096_kernel(
    A_ptr, B_ptr, C_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
    tl.static_assert(BLOCK_K % 16 == 0)
    tl.static_assert(BLOCK_M % 16 == 0)
    tl.static_assert(BLOCK_N % 16 == 0)

    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
    b_ptrs = B_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk

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

    k = 0
    while k < K:
        a = tl.load(
            a_ptrs,
            mask=(offs_m[:, None] < M) & (offs_k[None, :] + k < K),
            other=0.0,
        )
        b = tl.load(
            b_ptrs,
            mask=(offs_k[:, None] + k < K) & (offs_n[None, :] < N),
            other=0.0,
        )
        acc += tl.dot(a, b)
        k += BLOCK_K
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    c = acc.to(tl.float16)
    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(
        c_ptrs,
        c,
        mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
    )


def run(*args, **kwargs):
    A = None
    B = None
    if len(args) >= 1:
        A = args[0]
    if len(args) >= 2:
        B = args[1]
    if 'A' in kwargs:
        A = kwargs['A']
    if 'B' in kwargs:
        B = kwargs['B']

    if A is None or B is None:
        raise ValueError("run(A, B): both A and B must be provided")

    if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
        raise TypeError("A and B must be torch.Tensor")

    # Validate dtypes and shapes
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError("A and B must be float16 tensors")

    if A.ndim != 2 or B.ndim != 2:
        raise ValueError("A and B must be 2D tensors")

    M, K_a = A.shape
    N_b, K_b = B.shape

    N_SPEC = 28672
    K_SPEC = 4096

    if K_a != K_SPEC or K_b != K_SPEC:
        raise ValueError(f"K dimension must be {K_SPEC}; got A.shape[1]={K_a}, B.shape[1]={K_b}")
    if N_b != N_SPEC:
        raise ValueError(f"B.shape[0] (N) must be {N_SPEC}; got {N_b}")

    # Device management
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available; Triton kernel requires a CUDA-capable device")

    # Choose compute device
    compute_device = None
    if A.is_cuda:
        compute_device = A.device
    if B.is_cuda:
        # If both CUDA and different devices, prefer A's; else use B's
        compute_device = A.device if A.is_cuda else B.device
        if A.is_cuda and A.device != B.device:
            # Move B to A's device to compute
            pass

    if compute_device is None:
        compute_device = torch.device('cuda')

    # Move inputs to compute_device if needed
    if not A.is_cuda or A.device != compute_device:
        A_dev = A.to(device=compute_device, dtype=torch.float16, non_blocking=True)
    else:
        A_dev = A

    if not B.is_cuda or B.device != compute_device:
        B_dev = B.to(device=compute_device, dtype=torch.float16, non_blocking=True)
    else:
        B_dev = B

    # Prepare output on compute_device
    C_dev = torch.empty((M, N_SPEC), device=compute_device, dtype=torch.float16)

    # Strides (in elements)
    stride_am = A_dev.stride(0)
    stride_ak = A_dev.stride(1)
    stride_bn = B_dev.stride(0)
    stride_bk = B_dev.stride(1)
    stride_cm = C_dev.stride(0)
    stride_cn = C_dev.stride(1)

    # Grid
    def grid(meta):
        return (
            triton.cdiv(M, meta['BLOCK_M']),
            triton.cdiv(N_SPEC, meta['BLOCK_N']),
        )

    _gemm_n_28672_k_4096_kernel[grid](
        A_dev, B_dev, C_dev,
        M, N_SPEC, K_SPEC,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_cm, stride_cn,
    )

    # Move result back to original device if both inputs were originally on CPU
    # If any input was originally CUDA, return on that CUDA device (A's CUDA device takes precedence)
    if (not A.is_cuda) and (not B.is_cuda):
        return C_dev.cpu()
    else:
        # If A was originally CUDA and not on compute_device, move to A's original device?
        # Requirement: preserve original tensor devices and restore them for outputs.
        # Use A's original CUDA device if it was CUDA; else use B's original CUDA device.
        target_device = A.device if A.is_cuda else (B.device if B.is_cuda else compute_device)
        if C_dev.device != target_device:
            return C_dev.to(target_device, non_blocking=True)
        return C_dev
scrolls · 169 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON