Skip to content
KernelIndex
Search⌘K

gpt-5_triton_793693

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-793693?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 n6144 k4096fp16 · [1, 4096]
NVIDIA B200
88.2µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [15, 4096]
NVIDIA B200
89.0µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [48, 4096]
NVIDIA B200
89.1µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [8, 4096]
NVIDIA B200
89.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [2, 4096]
NVIDIA B200
89.3µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [4, 4096]
NVIDIA B200
89.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [35, 4096]
NVIDIA B200
89.3µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [24, 4096]
NVIDIA B200
89.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [32, 4096]
NVIDIA B200
89.3µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [16, 4096]
NVIDIA B200
89.4µs
#5 of 5
2025-10-16
Show all 43 measurements ›
GEMM n6144 k4096fp16 · [7, 4096]
NVIDIA B200
89.4µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [56, 4096]
NVIDIA B200
89.5µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [64, 4096]
NVIDIA B200
89.5µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [72, 4096]
NVIDIA B200
89.5µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [40, 4096]
NVIDIA B200
89.5µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [70, 4096]
NVIDIA B200
89.7µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [80, 4096]
NVIDIA B200
89.9µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [88, 4096]
NVIDIA B200
90.0µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [152, 4096]
NVIDIA B200
90.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [176, 4096]
NVIDIA B200
90.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [184, 4096]
NVIDIA B200
90.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [136, 4096]
NVIDIA B200
90.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [168, 4096]
NVIDIA B200
90.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [144, 4096]
NVIDIA B200
90.1µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [200, 4096]
NVIDIA B200
90.2µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [160, 4096]
NVIDIA B200
90.2µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [96, 4096]
NVIDIA B200
90.2µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [192, 4096]
NVIDIA B200
90.2µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [208, 4096]
NVIDIA B200
90.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [112, 4096]
NVIDIA B200
90.3µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [104, 4096]
NVIDIA B200
90.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [216, 4096]
NVIDIA B200
90.4µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [120, 4096]
NVIDIA B200
90.7µs
#5 of 5
2025-10-16
GEMM n6144 k4096fp16 · [224, 4096]
NVIDIA B200
91.0µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [232, 4096]
NVIDIA B200
91.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [128, 4096]
NVIDIA B200
91.5µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [240, 4096]
NVIDIA B200
91.9µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [248, 4096]
NVIDIA B200
92.1µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [256, 4096]
NVIDIA B200
92.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [972, 4096]
NVIDIA B200
173.0µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [2053, 4096]
NVIDIA B200
252.3µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [2379, 4096]
NVIDIA B200
320.7µs
#6 of 6
2025-10-16
GEMM n6144 k4096fp16 · [8192, 4096]
NVIDIA B200
939.3µs
#6 of 6
2025-10-16

Reported · How evidence levels are derived →

Source and license

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

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': 64},  num_stages=4, num_warps=8),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 128}, num_stages=5, num_warps=8),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 256, 'BLOCK_K': 64},  num_stages=4, num_warps=4),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64},  num_stages=4, num_warps=4),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 128, 'BLOCK_K': 64},  num_stages=4, num_warps=4),
    ],
    key=['M'],
)
@triton.jit
def _gemm_n_6144_k_4096_kernel(
    A_ptr, B_ptr, C_ptr,
    M, N, K,
    stride_am, stride_ak,    # A: [M, K]
    stride_bn, stride_bk,    # B: [N, K]
    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)

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

    k0 = 0
    while k0 < K:
        a_ptrs = A_ptr + (offs_m[:, None] * stride_am + (k0 + offs_k)[None, :] * stride_ak)
        b_ptrs = B_ptr + (offs_n[None, :] * stride_bn + (k0 + offs_k)[:, None] * stride_bk)
        a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & ((k0 + offs_k)[None, :] < K), other=0.0)
        b = tl.load(b_ptrs, mask=(offs_n[None, :] < N) & ((k0 + offs_k)[:, None] < K), other=0.0)
        acc += tl.dot(a, b)
        k0 += 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(*args, **kwargs):
    if len(args) == 2 and not kwargs:
        A, B = args
    else:
        A = kwargs.get('A', args[0] if len(args) > 0 else None)
        B = kwargs.get('B', args[1] if len(args) > 1 else None)
        if A is None or B is None:
            raise ValueError("run expects tensors A and B as positional or keyword arguments")

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

    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, KA = A.shape
    NB, KB = B.shape
    if KA != KB:
        raise ValueError(f"Incompatible inner dimensions: A is (*, {KA}), B is (*, {KB})")
    if NB != 6144:
        raise ValueError(f"B must have N=6144 as the first dimension, got {NB}")
    if KB != 4096:
        raise ValueError(f"B must have K=4096 as the second dimension, got {KB}")

    cuda_available = torch.cuda.is_available()
    A_is_cuda = A.is_cuda
    B_is_cuda = B.is_cuda

    if (A_is_cuda or B_is_cuda) and not cuda_available:
        raise RuntimeError("CUDA tensors provided but CUDA is not available")

    # Choose device: prefer GPU if available or if any input is on GPU
    if A_is_cuda:
        device = A.device
    elif B_is_cuda:
        device = B.device
    else:
        device = torch.device('cuda') if cuda_available else torch.device('cpu')

    if device.type == 'cpu' and not cuda_available:
        return torch.matmul(A, B.T)

    # Move to the chosen CUDA device if needed
    if device.type == 'cuda':
        dev_index = device.index if device.index is not None else 0
        A_dev = A.cuda(dev_index, non_blocking=True).contiguous()
        B_dev = B.cuda(dev_index, non_blocking=True).contiguous()
    else:
        A_dev = A.contiguous()
        B_dev = B.contiguous()

    M = A_dev.shape[0]
    K = A_dev.shape[1]
    N = B_dev.shape[0]

    C_dev = torch.empty((M, N), dtype=torch.float16, device=A_dev.device)

    stride_am, stride_ak = A_dev.stride()
    stride_bn, stride_bk = B_dev.stride()
    stride_cm, stride_cn = C_dev.stride()

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

    _gemm_n_6144_k_4096_kernel[grid](
        A_dev, B_dev, C_dev,
        M, N, K,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_cm, stride_cn,
    )

    # Move result back to A's original device
    if A.device == C_dev.device:
        return C_dev
    else:
        return C_dev.to(A.device, 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