Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritona796ab

gpt-o3_triton_a796ab · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-a796ab?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 n128 k2048fp16 · [1, 2048]
NVIDIA B200
51.4µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [8, 2048]
NVIDIA B200
52.5µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [4, 2048]
NVIDIA B200
52.6µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [2, 2048]
NVIDIA B200
53.2µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [5, 2048]
NVIDIA B200
53.2µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [16, 2048]
NVIDIA B200
53.4µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [6, 2048]
NVIDIA B200
53.4µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [34, 2048]
NVIDIA B200
54.0µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [17, 2048]
NVIDIA B200
54.8µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [32, 2048]
NVIDIA B200
55.0µs
#5 of 7
2025-10-16
Show all 25 measurements ›
GEMM n128 k2048fp16 · [25, 2048]
NVIDIA B200
55.3µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [64, 2048]
NVIDIA B200
55.3µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [172, 2048]
NVIDIA B200
55.4µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [63, 2048]
NVIDIA B200
56.1µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [128, 2048]
NVIDIA B200
57.3µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [492, 2048]
NVIDIA B200
58.0µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [289, 2048]
NVIDIA B200
58.9µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [8828, 2048]
NVIDIA B200
60.7µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [93, 2048]
NVIDIA B200
61.1µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [952, 2048]
NVIDIA B200
62.4µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [11006, 2048]
NVIDIA B200
65.0µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [12251, 2048]
NVIDIA B200
65.0µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [12853, 2048]
NVIDIA B200
65.1µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [14915, 2048]
NVIDIA B200
67.3µs
#5 of 7
2025-10-16
GEMM n128 k2048fp16 · [16294, 2048]
NVIDIA B200
67.6µs
#5 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:4e11feadfd6f90d6e9e5bde81706ab18677688caf7d4c78ec7b940c13937b104
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, tl.trans(b)) # (BLOCK_M, BLOCK_N)
num-warps = 8num_warps=8,
stages = 4num_stages=4,
tile-k = 32BLOCK_K = 32
tile-m = 64BLOCK_M = 64
tile-n = 128BLOCK_N = 128 # covers the whole N dimension

Kernel source

main.py180 lines
import math
from typing import Any, Dict, Tuple

import torch
import triton
import triton.language as tl


# -----------------------------------------------------------------------------
#                                TRITON KERNEL
# -----------------------------------------------------------------------------
@triton.jit
def _gemm_n128_k2048_kernel(
    A_ptr, B_ptr, C_ptr,
    M,                               # run–time size of the M dimension
    stride_am, stride_ak,            # strides for A  (row-major)
    stride_bn, stride_bk,            # strides for B  (row-major)
    stride_cm, stride_cn,            # strides for C  (row-major)
    BLOCK_M: tl.constexpr,           # tile sizes (compile–time constants)
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """
    Kernel computing C = A @ B.T
      A : [M, 2048]   (row-major, fp16)
      B : [128, 2048] (row-major, fp16) – accessed transposed
      C : [M, 128]    (row-major, fp16)

    Every program instance produces a tile of shape [BLOCK_M, BLOCK_N] of C.
    We split the workload only along the M dimension (N is fixed at 128).
    """
    # --------------------------------------------------------------------- #
    #  Identify the tile this program instance is responsible for           #
    # --------------------------------------------------------------------- #
    pid_m = tl.program_id(0)
    m_start = pid_m * BLOCK_M

    # Offsets inside the tile
    m_offsets = m_start + tl.arange(0, BLOCK_M)          # (BLOCK_M,)
    n_offsets = tl.arange(0, BLOCK_N)                    # (BLOCK_N,)

    # Accumulator – keep it in fp32 for accuracy
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # --------------------------------------------------------------------- #
    #  Iterate over K dimension (2048) in chunks of BLOCK_K                 #
    # --------------------------------------------------------------------- #
    K_ITERATIONS: tl.constexpr = 2048 // BLOCK_K
    for k_iter in tl.static_range(K_ITERATIONS):
        k_start   = k_iter * BLOCK_K
        k_offsets = k_start + tl.arange(0, BLOCK_K)      # (BLOCK_K,)

        # ----------------------------  Load A  --------------------------- #
        # A[m, k] with row-major stride
        a_ptrs = A_ptr + (m_offsets[:, None] * stride_am) + (k_offsets[None, :] * stride_ak)
        a = tl.load(
            a_ptrs,
            mask=m_offsets[:, None] < M,                 # guard rows that fall outside real M
            other=0.0
        )                                                # (BLOCK_M, BLOCK_K)

        # ----------------------------  Load B  --------------------------- #
        # B is stored [n, k] but we need B.T              -> shape (k, n)
        b_ptrs = B_ptr + (n_offsets[:, None] * stride_bn) + (k_offsets[None, :] * stride_bk)
        b = tl.load(b_ptrs)                              # (BLOCK_N, BLOCK_K)

        # --------------------------  Compute GEMM  ----------------------- #
        # b is (N, K).  We want (K, N) so transpose on-the-fly.
        acc += tl.dot(a, tl.trans(b))                    # (BLOCK_M, BLOCK_N)

    # --------------------------------------------------------------------- #
    #  Store the results                                                    #
    # --------------------------------------------------------------------- #
    c_ptrs = C_ptr + (m_offsets[:, None] * stride_cm) + (n_offsets[None, :] * stride_cn)
    tl.store(
        c_ptrs,
        acc.to(tl.float16),
        mask=m_offsets[:, None] < M
    )


# -----------------------------------------------------------------------------
#                             KERNEL LAUNCHER
# -----------------------------------------------------------------------------
def _launch_kernel(A_dev: torch.Tensor, B_dev: torch.Tensor) -> torch.Tensor:
    """
    Low-level helper that assumes both inputs live on the same CUDA device and
    are already contiguous and of dtype float16.  Returns C on that device.
    """
    if A_dev.dtype != torch.float16 or B_dev.dtype != torch.float16:
        raise TypeError("Both A and B must be float16 tensors")

    if A_dev.shape[1] != 2048:
        raise ValueError(f"A must have second dimension 2048, got {A_dev.shape}")
    if list(B_dev.shape) != [128, 2048]:
        raise ValueError(f"B must have shape [128, 2048], got {B_dev.shape}")

    # -----------  Tensor sizes & strides  -------------------------------- #
    M = A_dev.shape[0]

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

    C_dev = torch.empty((M, 128), dtype=torch.float16, device=A_dev.device)
    stride_cm, stride_cn = C_dev.stride()

    # -----------  Kernel configuration  ---------------------------------- #
    BLOCK_M = 64
    BLOCK_N = 128         # covers the whole N dimension
    BLOCK_K = 32

    grid = (triton.cdiv(M, BLOCK_M),)

    _gemm_n128_k2048_kernel[grid](
        A_dev, B_dev, C_dev,
        M,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_cm, stride_cn,
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
        num_warps=8,
        num_stages=4,
    )
    return C_dev


# -----------------------------------------------------------------------------
#                           PUBLIC ENTRY POINT
# -----------------------------------------------------------------------------
def run(*args: Any, **kwargs: Dict[str, Any]) -> torch.Tensor:
    """
    High-level helper replicating the reference API:

        C = run(A, B)               # positional
        C = run(A=A_tensor, B=B)    # keyword

    Handles device management:
      • Moves CPU tensors to GPU if necessary.
      • Ensures both inputs are on the same device.
      • Sends the result back to CPU if both inputs were on CPU.
    """
    # -------------  Retrieve A and B arguments  -------------------------- #
    if len(args) >= 2:
        A, B = args[:2]
    else:
        try:
            A = kwargs["A"]
            B = kwargs["B"]
        except KeyError as exc:
            raise ValueError("run expects tensors A and B either as positional "
                             "arguments or as keywords 'A' and 'B'") from exc

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required but not available")

    # -------------  Decide target CUDA device  --------------------------- #
    if A.is_cuda and B.is_cuda:
        target_device = A.device
        if B.device != target_device:
            raise RuntimeError("A and B must be on the same device")
    elif A.is_cuda:
        target_device = A.device
    elif B.is_cuda:
        target_device = B.device
    else:
        target_device = torch.device("cuda")

    # -------------  Move inputs to GPU & make contiguous ----------------- #
    A_dev = A.to(target_device, copy=False).contiguous()
    B_dev = B.to(target_device, copy=False).contiguous()

    # -------------  Launch the kernel ------------------------------------ #
    C_dev = _launch_kernel(A_dev, B_dev)

    # -------------  Move result back if inputs were on CPU --------------- #
    if (not A.is_cuda) and (not B.is_cuda):
        return C_dev.cpu()
    return C_dev
scrolls · 180 lines total

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

Best evidence level for this revision: reported

JSON