Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton93df2b

gpt-o3_triton_93df2b · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

29 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n2048 k4096fp16 · [1, 4096]
NVIDIA B200
144.6µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [2, 4096]
NVIDIA B200
145.2µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [4, 4096]
NVIDIA B200
148.0µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [5, 4096]
NVIDIA B200
153.7µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [6, 4096]
NVIDIA B200
156.4µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8, 4096]
NVIDIA B200
162.6µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15, 4096]
NVIDIA B200
178.5µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16, 4096]
NVIDIA B200
180.7µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [17, 4096]
NVIDIA B200
185.9µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [25, 4096]
NVIDIA B200
206.2µs
#5 of 7
2025-10-16
Show all 29 measurements ›
GEMM n2048 k4096fp16 · [32, 4096]
NVIDIA B200
219.8µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [34, 4096]
NVIDIA B200
229.2µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [63, 4096]
NVIDIA B200
299.5µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [64, 4096]
NVIDIA B200
301.4µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [93, 4096]
NVIDIA B200
366.4µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [172, 4096]
NVIDIA B200
421.3µs
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [289, 4096]
NVIDIA B200
421.9µs
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [492, 4096]
NVIDIA B200
422.3µs
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [952, 4096]
NVIDIA B200
422.7µs
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [969, 4096]
NVIDIA B200
422.9µs
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [128, 4096]
NVIDIA B200
442.7µs
#7 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8828, 4096]
NVIDIA B200
3.41ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11006, 4096]
NVIDIA B200
4.32ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11938, 4096]
NVIDIA B200
4.64ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12251, 4096]
NVIDIA B200
4.66ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12853, 4096]
NVIDIA B200
4.70ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [14915, 4096]
NVIDIA B200
5.51ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15813, 4096]
NVIDIA B200
5.93ms
#6 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16294, 4096]
NVIDIA B200
5.96ms
#6 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:659280037a402819bfb6259b3050f2343afa64265cef86a49d914a57ad704516
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) # (M,K) x (K,N) -> (M,N)
num-warps = 8num_warps=8,
stages = 4num_stages=4
tile-k = 64BLOCK_K = 64
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

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


###############################################################################
# Kernel
###############################################################################

@triton.jit
def _gemm_n2048_k4096_kernel(
    A_ptr, B_ptr, C_ptr,                                # pointers to matrices
    M: tl.int32,                                        # runtime M dimension
    stride_am: tl.int32, stride_ak: tl.int32,           # A strides
    stride_bn: tl.int32, stride_bk: tl.int32,           # B strides
    stride_cm: tl.int32, stride_cn: tl.int32,           # C strides
    BLOCK_M: tl.constexpr,                              # tile sizes
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """
    Compute C[M, 2048] = A[M, 4096] @ B[2048, 4096].T  (row–major tensors)
    Every program instance (CTA) computes a BLOCK_M x BLOCK_N tile of C.
    """

    # ---------------------- CTA indices -----------------------------
    pid_m = tl.program_id(0)        # block row index
    pid_n = tl.program_id(1)        # block col index

    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]

    # pointers for the tile of C that we will write
    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < 2048)

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

    # ---------------------- main k loop -----------------------------
    K_TOTAL = 4096
    for k in range(0, K_TOTAL, BLOCK_K):
        offs_k = k + tl.arange(0, BLOCK_K)                # [BLOCK_K]

        # ---- load A sub-tile : shape (BLOCK_M, BLOCK_K) -------
        a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
        a_mask = offs_m[:, None] < M                      # K dimension is always in range
        a = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)

        # ---- load B sub-tile (as KxN) : shape (BLOCK_K, BLOCK_N) -------
        b_ptrs = B_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk
        # offs_n < 2048 always by construction, offs_k < 4096 in loop bounds
        b = tl.load(b_ptrs).to(tl.float32)

        # ---- accumulate -------------------------------------------------
        acc += tl.dot(a, b)                               # (M,K) x (K,N) -> (M,N)

    # ---------------------- store results -------------------------------
    c = acc.to(tl.float16)
    tl.store(c_ptrs, c, mask=c_mask)


###############################################################################
# Host / wrapper
###############################################################################

def _check_inputs(A: torch.Tensor, B: torch.Tensor):
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError("A and B must be float16")
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError("A and B must be 2-D matrices")
    M, K_a = A.shape
    N, K_b = B.shape
    if N != 2048 or K_a != 4096 or K_b != 4096:
        raise ValueError(
            f"Expected A shape [M,4096], B shape [2048,4096], got {A.shape} and {B.shape}"
        )


def run(*args, **kwargs):
    """
    Entry-point that mimics the reference `run` interface:

    >>> C = run(A, B)

    Handles device placement automatically and works with positional
    or keyword arguments (A, B).
    """
    if len(args) + len(kwargs) != 2:
        raise TypeError("run expects exactly two tensor arguments, A and B")

    # Resolve arguments regardless of positional / keyword style
    if args:
        A = args[0]
        B = args[1] if len(args) > 1 else kwargs.get("B", None)
    else:
        A = kwargs.get("A", None)
        B = kwargs.get("B", None)

    if A is None or B is None:
        raise TypeError("Both A and B tensors must be provided")

    _check_inputs(A, B)

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device not available – Triton kernel cannot run")

    # Save original devices to restore later
    orig_device_A = A.device
    orig_device_B = B.device

    # Move to CUDA if necessary (non-blocking whenever possible)
    device = torch.device("cuda")
    A_cuda = A.to(device, non_blocking=True)
    B_cuda = B.to(device, non_blocking=True)

    M = A_cuda.shape[0]

    # Prepare output tensor
    C_cuda = torch.empty((M, 2048), dtype=torch.float16, device=device)

    # Strides (row-major tensors)
    stride_am, stride_ak = A_cuda.stride()
    stride_bn, stride_bk = B_cuda.stride()
    stride_cm, stride_cn = C_cuda.stride()

    # Launch configuration
    BLOCK_M = 128
    BLOCK_N = 128
    BLOCK_K = 64

    grid = (
        triton.cdiv(M, BLOCK_M),   # blocks along M
        2048 // BLOCK_N            # blocks along N (exact because 2048 % 128 == 0)
    )

    _gemm_n2048_k4096_kernel[grid](
        A_cuda, B_cuda, C_cuda,
        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
    )

    # Move back to original device of A (reference implementation behaviour)
    C_out = C_cuda.to(orig_device_A, non_blocking=True)
    return C_out
scrolls · 153 lines total

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

Best evidence level for this revision: reported

JSON