Skip to content
KernelIndex
Search⌘K

gpt-o3_triton_189201

gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

17 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n256 k7168fp16 · [1, 7168]
NVIDIA B200
290.3µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [4, 7168]
NVIDIA B200
290.3µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [901, 7168]
NVIDIA B200
292.0µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [16, 7168]
NVIDIA B200
292.7µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [15, 7168]
NVIDIA B200
296.8µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [14, 7168]
NVIDIA B200
297.4µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [11948, 7168]
NVIDIA B200
298.1µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [14104, 7168]
NVIDIA B200
299.2µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [54, 7168]
NVIDIA B200
302.7µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [32, 7168]
NVIDIA B200
303.2µs
#5 of 7
2025-10-16
Show all 17 measurements ›
GEMM n256 k7168fp16 · [57, 7168]
NVIDIA B200
304.1µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [53, 7168]
NVIDIA B200
304.5µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [58, 7168]
NVIDIA B200
306.5µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [55, 7168]
NVIDIA B200
307.3µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [56, 7168]
NVIDIA B200
308.2µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [63, 7168]
NVIDIA B200
310.2µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [80, 7168]
NVIDIA B200
315.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:b6694ec82d45fe33e3d4bdd9bbb95cacabc84dfb6bcc5db38fc2de95286dea1c
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) # (BM, BN)
num-warps = 8num_warps=8, num_stages=4,
stages = 4num_warps=8, num_stages=4,
tile-k = 128BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 32

Kernel source

main.py162 lines
import math
import torch
import triton
import triton.language as tl

# -----------------------------------------------------------------------------
# Problem constants (host-side)
# -----------------------------------------------------------------------------
K_TOTAL: int = 7_168          # fixed K dimension
N_CONST: int = 256            # fixed N dimension
DTYPE = torch.float16         # all tensors are fp16

# -----------------------------------------------------------------------------
# Triton kernel
# -----------------------------------------------------------------------------
@triton.jit
def _gemm_kernel(
    A_ptr, B_ptr, C_ptr,                       # pointers to tensors
    M,                                         # run-time length of the M axis
    stride_am, stride_ak,                      # strides for A
    stride_bk, stride_bn,                      # strides for B (row-major [N,K])
    stride_cm, stride_cn,                      # strides for C
    BLOCK_M: tl.constexpr,                    # tile sizes (compile-time)
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """
    Compute C = A @ B.T

    A : [M, K_TOTAL]
    B : [N_CONST, K_TOTAL]  (row-major / torch-contiguous)
    C : [M, N_CONST]

    The kernel spans the whole N axis (256) inside a single program instance
    and tiles the M axis with BLOCK_M.
    """

    # ------------------------------------------------------------------
    # Program-ID → position of the tile on the M axis
    # ------------------------------------------------------------------
    pid_m   = tl.program_id(0)
    m_start = pid_m * BLOCK_M

    # ------------------------------------------------------------------
    # Indices within the current tile
    # ------------------------------------------------------------------
    offs_m = m_start + tl.arange(0, BLOCK_M)         # [BM]
    offs_n = tl.arange(0, BLOCK_N)                    # [BN]  (0 … 255)
    offs_k = tl.arange(0, BLOCK_K)                    # [BK]

    # Accumulator in fp32 for improved precision
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # ------------------------------------------------------------------
    # Iterate over the K axis (compile-time constant length 7168)
    # ------------------------------------------------------------------
    for k_off in tl.static_range(0, 7_168, BLOCK_K):
        k_idx = k_off + offs_k                        # [BK]

        # Pointers for the current K-slice
        a_ptrs = A_ptr + (offs_m[:, None] * stride_am +
                          k_idx[None, :] * stride_ak)         # (BM, BK)
        b_ptrs = B_ptr + (k_idx[:, None] * stride_bk +
                          offs_n[None, :] * stride_bn)        # (BK, BN)

        # Guard against incomplete last M-tile
        mask_a = offs_m[:, None] < M

        a = tl.load(a_ptrs, mask=mask_a, other=0.0)           # (BM, BK)
        b = tl.load(b_ptrs)                                   # (BK, BN)

        acc += tl.dot(a, b)                                   # (BM, BN)

    # ------------------------------------------------------------------
    # Write results back to global memory
    # ------------------------------------------------------------------
    c = acc.to(tl.float16)                                    # fp32 → fp16
    c_ptrs = C_ptr + (offs_m[:, None] * stride_cm +
                      offs_n[None, :] * stride_cn)
    tl.store(c_ptrs, c, mask=offs_m[:, None] < M)


# -----------------------------------------------------------------------------
# Helper that launches the Triton kernel
# -----------------------------------------------------------------------------
def _launch_gemm(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required for Triton kernels but is not available.")

    # Move to GPU if necessary
    A_gpu = A.cuda() if A.device.type != "cuda" else A
    B_gpu = B.cuda() if B.device.type != "cuda" else B

    # Basic sanity checks
    if A_gpu.dtype != DTYPE or B_gpu.dtype != DTYPE:
        raise TypeError("A and B must be fp16 tensors.")
    if A_gpu.shape[1] != K_TOTAL:
        raise ValueError(f"A must have shape (M, {K_TOTAL})")
    if B_gpu.shape != (N_CONST, K_TOTAL):
        raise ValueError(f"B must have shape ({N_CONST}, {K_TOTAL})")

    # Ensure contiguous layout for coalesced accesses
    A_gpu = A_gpu.contiguous()
    B_gpu = B_gpu.contiguous()

    M = A_gpu.shape[0]
    C_gpu = torch.empty((M, N_CONST), dtype=DTYPE, device=A_gpu.device)

    # Kernel configuration
    BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 32
    grid = (triton.cdiv(M, BLOCK_M),)             # 1-D grid along M

    _gemm_kernel[grid](
        A_gpu, B_gpu, C_gpu,
        M,
        A_gpu.stride(0), A_gpu.stride(1),
        B_gpu.stride(1), B_gpu.stride(0),
        C_gpu.stride(0), C_gpu.stride(1),
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        num_warps=8, num_stages=4,
    )

    return C_gpu


# -----------------------------------------------------------------------------
# Public entry point (mirrors the reference interface)
# -----------------------------------------------------------------------------
def run(*args, **kwargs):
    """
    Usage:
        C = run(A, B)
        C = run(A=<tensor>, B=<tensor>)

    The result tensor is returned on the same device as A.
    """
    # ------------------------------------------------------ #
    # Parse arguments
    # ------------------------------------------------------ #
    if len(args) + len(kwargs) != 2:
        raise ValueError("run expects exactly two tensor arguments: A and B")

    if len(args) == 2:
        A, B = args
    else:
        A = kwargs.pop("A", None)
        B = kwargs.pop("B", None)
        if A is None or B is None:
            raise ValueError("Both A and B must be provided.")
        if kwargs:
            raise ValueError(f"Unexpected keyword arguments: {tuple(kwargs.keys())}")

    # ------------------------------------------------------ #
    # Execute kernel
    # ------------------------------------------------------ #
    C_gpu = _launch_gemm(A, B)

    # Return result on caller's original device
    return C_gpu.to(A.device)


__all__ = ["run"]
scrolls · 162 lines total

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

Best evidence level for this revision: reported

JSON