Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton63bfe6

gpt-o3_triton_63bfe6 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

25 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n5120 k2048fp16 · [289, 2048]
NVIDIA B200
61.8µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [172, 2048]
NVIDIA B200
64.1µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [1, 2048]
NVIDIA B200
64.6µs
#5 of 6
2025-10-16
GEMM n5120 k2048fp16 · [6, 2048]
NVIDIA B200
64.8µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8, 2048]
NVIDIA B200
64.8µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [5, 2048]
NVIDIA B200
65.1µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [4, 2048]
NVIDIA B200
65.2µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [2, 2048]
NVIDIA B200
65.3µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [25, 2048]
NVIDIA B200
65.4µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [16, 2048]
NVIDIA B200
65.4µs
#6 of 6
2025-10-16
Show all 25 measurements ›
GEMM n5120 k2048fp16 · [492, 2048]
NVIDIA B200
65.5µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [34, 2048]
NVIDIA B200
65.6µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [17, 2048]
NVIDIA B200
65.6µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [32, 2048]
NVIDIA B200
65.9µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [64, 2048]
NVIDIA B200
66.3µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [63, 2048]
NVIDIA B200
66.4µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [93, 2048]
NVIDIA B200
67.2µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [128, 2048]
NVIDIA B200
68.5µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [952, 2048]
NVIDIA B200
116.2µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8828, 2048]
NVIDIA B200
589.9µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [11006, 2048]
NVIDIA B200
751.1µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12251, 2048]
NVIDIA B200
776.1µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12853, 2048]
NVIDIA B200
878.3µs
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [14915, 2048]
NVIDIA B200
1.01ms
#6 of 6
2025-10-16
GEMM n5120 k2048fp16 · [16294, 2048]
NVIDIA B200
1.04ms
#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:a7572476136ff7667b6610eec154e63a44b7be17d38484adbf41a195fac9ddb2
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Techniques

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

mmaacc += tl.dot(a, b) # (BLOCK_M, BLOCK_N)
num-warps = 8num_warps = 8
stages = 4num_stages = 4
tile-k = 128BLOCK_M, BLOCK_N, BLOCK_K = 128, 128, 32

Kernel source

main.py152 lines
import math
from typing import Tuple

import torch
import triton
import triton.language as tl


################################################################################
#                              KERNEL                                           #
################################################################################
@triton.jit
def _gemm_n5120_k2048_kernel(
    A_ptr,                       # *fp16  [M, 2048]
    B_ptr,                       # *fp16  [5120, 2048]
    C_ptr,                       # *fp16  [M, 5120]
    M,                           # int32  dynamic
    stride_am, stride_ak,        # strides for A
    stride_bn, stride_bk,        # strides for B
    stride_cm, stride_cn,        # strides for C
    BLOCK_M: tl.constexpr,       # tile sizes
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """
    Compute C[M,5120] = A[M,2048] @ B[5120,2048]^T   (fp16 accumulate in fp32)
    The K dimension (2048) and N dimension (5120) are compile-time constants,
    which enables full loop unrolling and constant-folding in Triton.
    """
    # ------------------------------------------------------------------
    # Pointer arithmetic helpers
    # ------------------------------------------------------------------
    pid_m = tl.program_id(0)                   # program id along M dimension
    pid_n = tl.program_id(1)                   # program id along N dimension

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

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

    # Loop over K dimension – 2048 is constant, so we can completely unroll
    for k_iter in tl.static_range(0, 2048, BLOCK_K):
        k_curr = k_iter + offs_k                                            # (BLOCK_K,)

        # -----------------  Load A tile:  [BLOCK_M, BLOCK_K] ---------------
        a_ptrs = A_ptr + (offs_m[:, None] * stride_am) + (k_curr[None, :] * stride_ak)
        a_mask = (offs_m[:, None] < M) & (k_curr[None, :] < 2048)
        a = tl.load(a_ptrs, mask=a_mask, other=0.0)

        # -----------------  Load B^T tile:  [BLOCK_K, BLOCK_N] -------------
        # B is stored as (N, K); to access B^T we index as (k, n)
        b_ptrs = B_ptr + (offs_n[None, :] * stride_bn) + (k_curr[:, None] * stride_bk)
        b_mask = (offs_n[None, :] < 5120) & (k_curr[:, None] < 2048)
        b = tl.load(b_ptrs, mask=b_mask, other=0.0)

        # -----------------  Matrix multiply -------------------------------
        acc += tl.dot(a, b)   # (BLOCK_M, BLOCK_N)

    # ------------------------- Store results ------------------------------
    c_ptrs = C_ptr + (offs_m[:, None] * stride_cm) + (offs_n[None, :] * stride_cn)
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < 5120)
    tl.store(c_ptrs, acc.to(tl.float16), mask=c_mask)


################################################################################
#                              HOST / WRAPPER                                  #
################################################################################
def _assert_shapes(a: torch.Tensor, b: torch.Tensor) -> None:
    if a.ndim != 2 or b.ndim != 2:
        raise ValueError("A and B must be 2-D tensors.")
    if a.shape[1] != 2048:
        raise ValueError(f"A must have shape [M, 2048], got {a.shape}.")
    if b.shape != (5120, 2048):
        raise ValueError(f"B must have shape [5120, 2048], got {b.shape}.")
    if a.dtype != torch.float16 or b.dtype != torch.float16:
        raise ValueError("Both A and B must be of dtype torch.float16.")


def _pick_device(a: torch.Tensor, b: torch.Tensor) -> torch.device:
    """
    Decide which CUDA device to run on.
    Priority:
      1. If A is CUDA → use A.device
      2. Else if B is CUDA → use B.device
      3. Else → default CUDA device (id 0)
    """
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required for this Triton kernel but is not available.")
    if a.is_cuda:
        return a.device
    if b.is_cuda:
        return b.device
    return torch.device("cuda")    # default device (usually id 0)


def run(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    """
    Host-side convenience wrapper.
    • Moves inputs to GPU (B200) if necessary
    • Launches Triton GEMM kernel
    • Moves output back to the original devices of inputs
    """
    # --------------------- Argument checking ------------------------------
    _assert_shapes(A, B)
    compute_device = _pick_device(A, B)

    # Keep track of where we must put the answer
    want_cpu_output = (not A.is_cuda) and (not B.is_cuda)
    output_device = A.device if A.is_cuda else (B.device if B.is_cuda else torch.device("cpu"))

    # --------------------- Move inputs to correct device ------------------
    A_dev = A.to(compute_device, non_blocking=True) if not A.device == compute_device else A
    B_dev = B.to(compute_device, non_blocking=True) if not B.device == compute_device else B

    M = A_dev.shape[0]                        # dynamic dimension
    N = 5120                                  # constant
    K = 2048                                  # constant

    # Output tensor
    C_dev = torch.empty((M, N), dtype=torch.float16, device=compute_device)

    # --------------------- Kernel launch configuration --------------------
    BLOCK_M, BLOCK_N, BLOCK_K = 128, 128, 32
    num_warps = 8
    num_stages = 4

    grid: Tuple[int, int] = (
        triton.cdiv(M, BLOCK_M),
        triton.cdiv(N, BLOCK_N),
    )

    _gemm_n5120_k2048_kernel[grid](
        A_dev, B_dev, C_dev,
        M,
        A_dev.stride(0), A_dev.stride(1),
        B_dev.stride(0), B_dev.stride(1),
        C_dev.stride(0), C_dev.stride(1),
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
        num_warps=num_warps,
        num_stages=num_stages,
    )

    # --------------------- Return to original device ----------------------
    if want_cpu_output:
        return C_dev.cpu()
    if C_dev.device != output_device:
        return C_dev.to(output_device, non_blocking=True)
    return C_dev
scrolls · 152 lines total

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

Best evidence level for this revision: reported

JSON