gpt-o3 / triton63bfe6
gpt-o3_triton_63bfe6 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 152 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-63bfe6?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
Show all 25 measurements ›Showing all 25 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:a7572476136ff7667b6610eec154e63a44b7be17d38484adbf41a195fac9ddb2
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.
mma
acc += tl.dot(a, b) # (BLOCK_M, BLOCK_N)num-warps = 8
num_warps = 8stages = 4
num_stages = 4tile-k = 128
BLOCK_M, BLOCK_N, BLOCK_K = 128, 128, 32Kernel source
main.py152 lines
import math
from typing import Tuple
import torch
import triton
import triton.language as tl
################################################################################
# KERNEL #
################################################################################
@triton.jit
def _gemm_n5120_k2048_kernel(
A_ptr, # *fp16 [M, 2048]
B_ptr, # *fp16 [5120, 2048]
C_ptr, # *fp16 [M, 5120]
M, # int32 dynamic
stride_am, stride_ak, # strides for A
stride_bn, stride_bk, # strides for B
stride_cm, stride_cn, # strides for C
BLOCK_M: tl.constexpr, # tile sizes
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""
Compute C[M,5120] = A[M,2048] @ B[5120,2048]^T (fp16 accumulate in fp32)
The K dimension (2048) and N dimension (5120) are compile-time constants,
which enables full loop unrolling and constant-folding in Triton.
"""
# ------------------------------------------------------------------
# Pointer arithmetic helpers
# ------------------------------------------------------------------
pid_m = tl.program_id(0) # program id along M dimension
pid_n = tl.program_id(1) # program id along 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,)
# Accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K dimension – 2048 is constant, so we can completely unroll
for k_iter in tl.static_range(0, 2048, BLOCK_K):
k_curr = k_iter + offs_k # (BLOCK_K,)
# ----------------- Load A tile: [BLOCK_M, BLOCK_K] ---------------
a_ptrs = A_ptr + (offs_m[:, None] * stride_am) + (k_curr[None, :] * stride_ak)
a_mask = (offs_m[:, None] < M) & (k_curr[None, :] < 2048)
a = tl.load(a_ptrs, mask=a_mask, other=0.0)
# ----------------- Load B^T tile: [BLOCK_K, BLOCK_N] -------------
# B is stored as (N, K); to access B^T we index as (k, n)
b_ptrs = B_ptr + (offs_n[None, :] * stride_bn) + (k_curr[:, None] * stride_bk)
b_mask = (offs_n[None, :] < 5120) & (k_curr[:, None] < 2048)
b = tl.load(b_ptrs, mask=b_mask, other=0.0)
# ----------------- Matrix multiply -------------------------------
acc += tl.dot(a, b) # (BLOCK_M, BLOCK_N)
# ------------------------- Store results ------------------------------
c_ptrs = C_ptr + (offs_m[:, None] * stride_cm) + (offs_n[None, :] * stride_cn)
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < 5120)
tl.store(c_ptrs, acc.to(tl.float16), mask=c_mask)
################################################################################
# HOST / WRAPPER #
################################################################################
def _assert_shapes(a: torch.Tensor, b: torch.Tensor) -> None:
if a.ndim != 2 or b.ndim != 2:
raise ValueError("A and B must be 2-D tensors.")
if a.shape[1] != 2048:
raise ValueError(f"A must have shape [M, 2048], got {a.shape}.")
if b.shape != (5120, 2048):
raise ValueError(f"B must have shape [5120, 2048], got {b.shape}.")
if a.dtype != torch.float16 or b.dtype != torch.float16:
raise ValueError("Both A and B must be of dtype torch.float16.")
def _pick_device(a: torch.Tensor, b: torch.Tensor) -> torch.device:
"""
Decide which CUDA device to run on.
Priority:
1. If A is CUDA → use A.device
2. Else if B is CUDA → use B.device
3. Else → default CUDA device (id 0)
"""
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this Triton kernel but is not available.")
if a.is_cuda:
return a.device
if b.is_cuda:
return b.device
return torch.device("cuda") # default device (usually id 0)
def run(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Host-side convenience wrapper.
• Moves inputs to GPU (B200) if necessary
• Launches Triton GEMM kernel
• Moves output back to the original devices of inputs
"""
# --------------------- Argument checking ------------------------------
_assert_shapes(A, B)
compute_device = _pick_device(A, B)
# Keep track of where we must put the answer
want_cpu_output = (not A.is_cuda) and (not B.is_cuda)
output_device = A.device if A.is_cuda else (B.device if B.is_cuda else torch.device("cpu"))
# --------------------- Move inputs to correct device ------------------
A_dev = A.to(compute_device, non_blocking=True) if not A.device == compute_device else A
B_dev = B.to(compute_device, non_blocking=True) if not B.device == compute_device else B
M = A_dev.shape[0] # dynamic dimension
N = 5120 # constant
K = 2048 # constant
# Output tensor
C_dev = torch.empty((M, N), dtype=torch.float16, device=compute_device)
# --------------------- Kernel launch configuration --------------------
BLOCK_M, BLOCK_N, BLOCK_K = 128, 128, 32
num_warps = 8
num_stages = 4
grid: Tuple[int, int] = (
triton.cdiv(M, BLOCK_M),
triton.cdiv(N, BLOCK_N),
)
_gemm_n5120_k2048_kernel[grid](
A_dev, B_dev, C_dev,
M,
A_dev.stride(0), A_dev.stride(1),
B_dev.stride(0), B_dev.stride(1),
C_dev.stride(0), C_dev.stride(1),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
num_warps=num_warps,
num_stages=num_stages,
)
# --------------------- Return to original device ----------------------
if want_cpu_output:
return C_dev.cpu()
if C_dev.device != output_device:
return C_dev.to(output_device, non_blocking=True)
return C_devscrolls · 152 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON