Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton4c9c32

gpt-o3_triton_4c9c32 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-4c9c32?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 n28672 k4096fp16 · [80, 4096]
NVIDIA B200
85.3µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [64, 4096]
NVIDIA B200
85.4µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [70, 4096]
NVIDIA B200
85.5µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [96, 4096]
NVIDIA B200
85.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [72, 4096]
NVIDIA B200
85.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [48, 4096]
NVIDIA B200
86.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [88, 4096]
NVIDIA B200
86.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [56, 4096]
NVIDIA B200
86.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [112, 4096]
NVIDIA B200
86.3µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [104, 4096]
NVIDIA B200
86.5µs
#6 of 8
2025-10-16
Show all 43 measurements ›
GEMM n28672 k4096fp16 · [32, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [4, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [15, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [16, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [35, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [1, 4096]
NVIDIA B200
86.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [40, 4096]
NVIDIA B200
86.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [128, 4096]
NVIDIA B200
86.9µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [24, 4096]
NVIDIA B200
86.9µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8, 4096]
NVIDIA B200
87.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [7, 4096]
NVIDIA B200
87.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2, 4096]
NVIDIA B200
87.1µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [120, 4096]
NVIDIA B200
87.2µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [208, 4096]
NVIDIA B200
124.5µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [216, 4096]
NVIDIA B200
124.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [192, 4096]
NVIDIA B200
125.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [224, 4096]
NVIDIA B200
125.1µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [136, 4096]
NVIDIA B200
125.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [232, 4096]
NVIDIA B200
125.7µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [240, 4096]
NVIDIA B200
125.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [184, 4096]
NVIDIA B200
125.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [152, 4096]
NVIDIA B200
125.9µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [168, 4096]
NVIDIA B200
125.9µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [160, 4096]
NVIDIA B200
126.0µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [200, 4096]
NVIDIA B200
126.1µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [176, 4096]
NVIDIA B200
126.3µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [248, 4096]
NVIDIA B200
126.4µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [256, 4096]
NVIDIA B200
126.7µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [144, 4096]
NVIDIA B200
127.2µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [972, 4096]
NVIDIA B200
322.1µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2053, 4096]
NVIDIA B200
609.0µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2379, 4096]
NVIDIA B200
685.4µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8192, 4096]
NVIDIA B200
2.41ms
#5 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:ff556ff243a36e31f4f56c6d97cf2d7522539056c7d8008d036404779d04992b
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)
num-warps = 8num_warps=8,
stages = 4num_stages=4
tile-k = 32BLOCK_K = 32
tile-m = 128BLOCK_M = 128
tile-n = 256BLOCK_N = 256

Kernel source

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


@triton.jit
def gemm_kernel(
    A_ptr, B_ptr, C_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    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]
    offs_k = tl.arange(0, BLOCK_K)                                # [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_n[None, :] * stride_bn +
                      offs_k[:, None] * stride_bk)                # [BLOCK_K, BLOCK_N]

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

    num_k_iters = tl.cdiv(K, BLOCK_K)
    for _ in range(num_k_iters):
        a = tl.load(
            a_ptrs,
            mask=(offs_m[:, None] < M) & (offs_k[None, :] < K),
            other=0.0
        )
        b = tl.load(
            b_ptrs,
            mask=(offs_n[None, :] < N) & (offs_k[:, None] < K),
            other=0.0
        )
        acc += tl.dot(a, b)

        offs_k += BLOCK_K
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    c_ptrs = C_ptr + (offs_m[:, None] * stride_cm +
                      offs_n[None, :] * stride_cn)
    acc = acc.to(tl.float16)
    tl.store(
        c_ptrs,
        acc,
        mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)
    )


def run(A: torch.Tensor, B: torch.Tensor):
    """
    High-performance GEMM on B200 GPUs.
    C = A @ B.T
    Shapes:
        A: [M, 4096]   (float16)
        B: [28672, 4096] (float16)
        C: [M, 28672]  (float16)
    """
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError("A and B must be 2-D tensors")
    if A.shape[1] != 4096 or B.shape[1] != 4096 or B.shape[0] != 28672:
        raise ValueError("Expected shapes: A [M, 4096], B [28672, 4096]")
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError("A and B must be float16")

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

    orig_device_A = A.device
    orig_device_B = B.device

    A_cuda = A.cuda() if not A.is_cuda else A
    B_cuda = B.cuda() if not B.is_cuda else B

    M = A_cuda.shape[0]
    N = 28672
    K = 4096

    C_cuda = torch.empty((M, N), device=A_cuda.device, dtype=torch.float16)

    BLOCK_M = 128
    BLOCK_N = 256
    BLOCK_K = 32

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

    gemm_kernel[grid](
        A_cuda, B_cuda, C_cuda,
        M, N, K,
        A_cuda.stride(0), A_cuda.stride(1),
        B_cuda.stride(0), B_cuda.stride(1),
        C_cuda.stride(0), C_cuda.stride(1),
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        num_warps=8,
        num_stages=4
    )

    torch.cuda.synchronize()

    if orig_device_A.type == "cuda":
        return C_cuda
    return C_cuda.cpu()


# Allow module import without immediate execution
if __name__ == "__main__":
    # Simple correctness test
    M_test = 256
    A_test = torch.randn((M_test, 4096), dtype=torch.float16)
    B_test = torch.randn((28672, 4096), dtype=torch.float16)
    C_ref = (A_test.float() @ B_test.t().float()).half()
    C_out = run(A_test, B_test)
    assert torch.allclose(C_ref, C_out, atol=1e-2, rtol=1e-2)
    print("Test passed!")
scrolls · 124 lines total

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

Best evidence level for this revision: reported

JSON