Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonef0d11

gpt-o3_triton_ef0d11 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-ef0d11?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 n4096 k4096fp16 · [2, 4096]
NVIDIA B200
103.6µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [1, 4096]
NVIDIA B200
103.6µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [4, 4096]
NVIDIA B200
104.1µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [7, 4096]
NVIDIA B200
105.2µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [40, 4096]
NVIDIA B200
105.3µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [8, 4096]
NVIDIA B200
105.4µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [16, 4096]
NVIDIA B200
105.4µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [15, 4096]
NVIDIA B200
105.4µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [32, 4096]
NVIDIA B200
105.4µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [24, 4096]
NVIDIA B200
105.5µs
#6 of 8
2025-10-16
Show all 43 measurements ›
GEMM n4096 k4096fp16 · [56, 4096]
NVIDIA B200
105.5µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [35, 4096]
NVIDIA B200
105.5µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [48, 4096]
NVIDIA B200
105.5µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [64, 4096]
NVIDIA B200
106.5µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [72, 4096]
NVIDIA B200
106.8µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [70, 4096]
NVIDIA B200
107.3µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [80, 4096]
NVIDIA B200
107.4µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [88, 4096]
NVIDIA B200
107.5µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [96, 4096]
NVIDIA B200
107.8µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [104, 4096]
NVIDIA B200
108.1µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [120, 4096]
NVIDIA B200
109.3µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [112, 4096]
NVIDIA B200
109.5µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [184, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [152, 4096]
NVIDIA B200
109.6µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [128, 4096]
NVIDIA B200
109.6µs
#7 of 9
2025-10-16
GEMM n4096 k4096fp16 · [168, 4096]
NVIDIA B200
109.6µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [200, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [256, 4096]
NVIDIA B200
109.6µs
#7 of 9
2025-10-16
GEMM n4096 k4096fp16 · [216, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [144, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [136, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [232, 4096]
NVIDIA B200
109.6µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [160, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [176, 4096]
NVIDIA B200
109.6µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [240, 4096]
NVIDIA B200
109.6µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [248, 4096]
NVIDIA B200
109.6µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [192, 4096]
NVIDIA B200
109.6µs
#7 of 8
2025-10-16
GEMM n4096 k4096fp16 · [224, 4096]
NVIDIA B200
109.6µs
#6 of 8
2025-10-16
GEMM n4096 k4096fp16 · [208, 4096]
NVIDIA B200
109.6µs
#7 of 9
2025-10-16
GEMM n4096 k4096fp16 · [972, 4096]
NVIDIA B200
120.0µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [2053, 4096]
NVIDIA B200
158.5µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [2379, 4096]
NVIDIA B200
181.3µs
#6 of 7
2025-10-16
GEMM n4096 k4096fp16 · [8192, 4096]
NVIDIA B200
384.0µs
#6 of 8
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:5d846b9f7c2d09afe4cefb4461449998b40106ce852f5d4cae2d736dd2cb3728
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) # fp32 accumulation
num-warps = 8num_warps=8,
stages = 4num_stages=4
tile-k = 32def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):
tile-m = 128def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):
tile-n = 128def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):

Kernel source

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


@triton.jit
def _gemm_4096_kernel(
    A_ptr, B_ptr, C_ptr,
    M,                                   # total rows of A / C
    stride_am, stride_ak,                # strides for A
    stride_bk, stride_bn,                # strides for B (transposed view K x N)
    stride_cm, stride_cn,                # strides for C
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr
):
    """
    Kernel computes C = A @ B, where
    A: [M, K]  (row-major)
    B: [K, N]  (row-major, supplied as B.T from host for contiguous K dimension)
    K = 4096, N = 4096 (constants)
    """
    pid_m = tl.program_id(0)                   # program id for M dimension
    pid_n = tl.program_id(1)                   # program id for 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]

    # Pointer to C
    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn

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

    # Loop over K axis
    for k0 in range(0, 4096, BLOCK_K):
        offs_k = k0 + tl.arange(0, BLOCK_K)

        a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak       # (BLOCK_M, BLOCK_K)
        b_ptrs = B_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn       # (BLOCK_K, BLOCK_N)

        a = tl.load(a_ptrs, mask=offs_m[:, None] < M, other=0.0)                         # mask for last incomplete block
        b = tl.load(b_ptrs)                                                              # K and N are constants, no mask

        acc += tl.dot(a, b)                                                              # fp32 accumulation

    # Write back
    c = acc.to(tl.float16)
    tl.store(c_ptrs, c, mask=offs_m[:, None] < M)


def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):
    # Grid dimensions
    M = A.shape[0]
    grid_m = triton.cdiv(M, BLOCK_M)
    grid_n = 32                               # 4096 / 128

    _gemm_4096_kernel[(grid_m, grid_n)](
        A, B_T, C,
        M,
        A.stride(0), A.stride(1),
        B_T.stride(0), B_T.stride(1),
        C.stride(0), C.stride(1),
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
        num_warps=8,
        num_stages=4
    )


def run(A: torch.Tensor, B: torch.Tensor):
    """
    Entry point that matches reference semantics:
    C = A @ B.T
    Handles device placement transparently.
    """

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernel")

    # Preserve original devices
    device_a = A.device
    device_b = B.device

    # Move to GPU if necessary
    A_gpu = A.cuda() if not A.is_cuda else A
    B_gpu = B.cuda() if not B.is_cuda else B

    # Materialize B^T as contiguous KxN tensor
    B_T = B_gpu.t().contiguous()

    # Allocate output tensor on GPU
    M = A_gpu.shape[0]
    C_gpu = torch.empty((M, 4096), device=A_gpu.device, dtype=torch.float16)

    # Launch Triton kernel
    _launch_kernel(A_gpu, B_T, C_gpu)

    # Move result back to the device of A (arbitrary choice if A & B differ)
    C_out = C_gpu.to(device_a)

    return C_out


# If this file is executed directly, run a quick correctness test
if __name__ == "__main__":
    torch.manual_seed(0)
    M_test = 512
    A_test = torch.randn((M_test, 4096), dtype=torch.float16)
    B_test = torch.randn((4096, 4096), dtype=torch.float16)

    C_ref = torch.matmul(A_test, B_test.t())
    C_triton = run(A_test, B_test)

    max_err = (C_ref - C_triton).abs().max()
    print("Max error:", max_err.item())
scrolls · 118 lines total

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

Best evidence level for this revision: reported

JSON