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
Show all 29 measurements ›Showing all 29 measurements ⌄
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.
mma
acc += tl.dot(a, b) # (M,K) x (K,N) -> (M,N)num-warps = 8
num_warps=8,stages = 4
num_stages=4tile-k = 64
BLOCK_K = 64tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128Kernel 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_outscrolls · 153 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON