Skip to content
KernelIndex
Search⌘K

gpt-5 / triton1bd4a7

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-1bd4a7?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 n4096 k14336fp16 · [15, 14336]
NVIDIA B200
1.06ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [136, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [144, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [16, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [8, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [152, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [24, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [40, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [32, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [35, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
Show all 43 measurements ›
GEMM n4096 k14336fp16 · [48, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [160, 14336]
NVIDIA B200
1.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [56, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [972, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [64, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [168, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [72, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [70, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [80, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [176, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [88, 14336]
NVIDIA B200
1.08ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [96, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [7, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [184, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [192, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [4, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [200, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [208, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [216, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [224, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [232, 14336]
NVIDIA B200
1.09ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [240, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [248, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [256, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [104, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [112, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [120, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [128, 14336]
NVIDIA B200
1.10ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [1, 14336]
NVIDIA B200
1.11ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [2, 14336]
NVIDIA B200
1.12ms
#6 of 6
2025-10-20
GEMM n4096 k14336fp16 · [2053, 14336]
NVIDIA B200
2.07ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [2379, 14336]
NVIDIA B200
3.16ms
#5 of 6
2025-10-20
GEMM n4096 k14336fp16 · [8192, 14336]
NVIDIA B200
7.45ms
#5 of 6
2025-10-20

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:8cd8cac18df360a37c9fc4b35cef81e63d0edf283dc5d5e973979dd667b60a00
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': 128}, num_warps=8, num_stages=5),
stages = 5triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 128}, num_warps=8, num_stages=5),

Kernel source

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


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 128}, num_warps=8, num_stages=5),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),
        triton.Config({'BLOCK_M': 256, 'BLOCK_N': 128, 'BLOCK_K': 128}, num_warps=8, num_stages=5),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 256, 'BLOCK_K': 128}, num_warps=4, num_stages=5),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 128}, num_warps=4, num_stages=5),
        triton.Config({'BLOCK_M': 256, 'BLOCK_N': 256, 'BLOCK_K': 64},  num_warps=16, num_stages=4),
    ],
    key=['M']
)
@triton.jit
def _gemm_mk_kn_to_mn_kernel(
    A_ptr, B_ptr, C_ptr,
    M, N, K,
    stride_am, stride_ak,   # A: [M, K]
    stride_bn, stride_bk,   # B: [N, K] but loaded as [K, N] via strides
    stride_cm, stride_cn,   # C: [M, N]
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
    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_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)

    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) & (k + offs_k[None, :] < K),
            other=0.0
        )
        b = tl.load(
            b_ptrs,
            mask=(k + offs_k[:, None] < K) & (offs_n[None, :] < N),
            other=0.0
        )
        acc += tl.dot(a, b)

        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk
        k += BLOCK_K

    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(A, B):
    if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
        raise TypeError("Inputs A and B must be torch.Tensor")

    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

    # Constants from specification
    REQUIRED_N = 4096
    REQUIRED_K = 14336

    if K_a != REQUIRED_K or K_b != REQUIRED_K:
        raise ValueError(f"K must be {REQUIRED_K}. Got A.shape[1]={K_a}, B.shape[1]={K_b}")
    if N_b != REQUIRED_N:
        raise ValueError(f"N must be {REQUIRED_N}. Got B.shape[0]={N_b}")

    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError("A and B must be of dtype torch.float16")

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run this Triton kernel, but torch.cuda.is_available() is False.")

    # Pick a CUDA device
    if A.is_cuda:
        cuda_dev = A.device
    elif B.is_cuda:
        cuda_dev = B.device
    else:
        cuda_dev = torch.device('cuda')

    # Preserve original devices without modifying inputs
    dev_A_orig = A.device
    dev_B_orig = B.device

    # Move to chosen CUDA device if needed
    A_gpu = A.to(device=cuda_dev, non_blocking=True) if A.device != cuda_dev else A
    B_gpu = B.to(device=cuda_dev, non_blocking=True) if B.device != cuda_dev else B

    # Shapes
    M = A_gpu.shape[0]
    N = B_gpu.shape[0]
    K = A_gpu.shape[1]

    # Allocate output on GPU
    C_gpu = torch.empty((M, N), dtype=torch.float16, device=cuda_dev)

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

    # Launch kernel
    _gemm_mk_kn_to_mn_kernel[grid](
        A_gpu, B_gpu, C_gpu,
        M, N, K,
        A_gpu.stride(0), A_gpu.stride(1),
        B_gpu.stride(0), B_gpu.stride(1),
        C_gpu.stride(0), C_gpu.stride(1),
    )

    # Move result back to the original device of A
    return C_gpu.to(dev_A_orig, non_blocking=True)
scrolls · 127 lines total

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

Best evidence level for this revision: reported

JSON