gpt-o3 / tritona796ab
gpt-o3_triton_a796ab · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 180 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-a796ab?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:4e11feadfd6f90d6e9e5bde81706ab18677688caf7d4c78ec7b940c13937b104
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, tl.trans(b)) # (BLOCK_M, BLOCK_N)num-warps = 8
num_warps=8,stages = 4
num_stages=4,tile-k = 32
BLOCK_K = 32tile-m = 64
BLOCK_M = 64tile-n = 128
BLOCK_N = 128 # covers the whole N dimensionKernel source
main.py180 lines
import math
from typing import Any, Dict, Tuple
import torch
import triton
import triton.language as tl
# -----------------------------------------------------------------------------
# TRITON KERNEL
# -----------------------------------------------------------------------------
@triton.jit
def _gemm_n128_k2048_kernel(
A_ptr, B_ptr, C_ptr,
M, # run–time size of the M dimension
stride_am, stride_ak, # strides for A (row-major)
stride_bn, stride_bk, # strides for B (row-major)
stride_cm, stride_cn, # strides for C (row-major)
BLOCK_M: tl.constexpr, # tile sizes (compile–time constants)
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""
Kernel computing C = A @ B.T
A : [M, 2048] (row-major, fp16)
B : [128, 2048] (row-major, fp16) – accessed transposed
C : [M, 128] (row-major, fp16)
Every program instance produces a tile of shape [BLOCK_M, BLOCK_N] of C.
We split the workload only along the M dimension (N is fixed at 128).
"""
# --------------------------------------------------------------------- #
# Identify the tile this program instance is responsible for #
# --------------------------------------------------------------------- #
pid_m = tl.program_id(0)
m_start = pid_m * BLOCK_M
# Offsets inside the tile
m_offsets = m_start + tl.arange(0, BLOCK_M) # (BLOCK_M,)
n_offsets = tl.arange(0, BLOCK_N) # (BLOCK_N,)
# Accumulator – keep it in fp32 for accuracy
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# --------------------------------------------------------------------- #
# Iterate over K dimension (2048) in chunks of BLOCK_K #
# --------------------------------------------------------------------- #
K_ITERATIONS: tl.constexpr = 2048 // BLOCK_K
for k_iter in tl.static_range(K_ITERATIONS):
k_start = k_iter * BLOCK_K
k_offsets = k_start + tl.arange(0, BLOCK_K) # (BLOCK_K,)
# ---------------------------- Load A --------------------------- #
# A[m, k] with row-major stride
a_ptrs = A_ptr + (m_offsets[:, None] * stride_am) + (k_offsets[None, :] * stride_ak)
a = tl.load(
a_ptrs,
mask=m_offsets[:, None] < M, # guard rows that fall outside real M
other=0.0
) # (BLOCK_M, BLOCK_K)
# ---------------------------- Load B --------------------------- #
# B is stored [n, k] but we need B.T -> shape (k, n)
b_ptrs = B_ptr + (n_offsets[:, None] * stride_bn) + (k_offsets[None, :] * stride_bk)
b = tl.load(b_ptrs) # (BLOCK_N, BLOCK_K)
# -------------------------- Compute GEMM ----------------------- #
# b is (N, K). We want (K, N) so transpose on-the-fly.
acc += tl.dot(a, tl.trans(b)) # (BLOCK_M, BLOCK_N)
# --------------------------------------------------------------------- #
# Store the results #
# --------------------------------------------------------------------- #
c_ptrs = C_ptr + (m_offsets[:, None] * stride_cm) + (n_offsets[None, :] * stride_cn)
tl.store(
c_ptrs,
acc.to(tl.float16),
mask=m_offsets[:, None] < M
)
# -----------------------------------------------------------------------------
# KERNEL LAUNCHER
# -----------------------------------------------------------------------------
def _launch_kernel(A_dev: torch.Tensor, B_dev: torch.Tensor) -> torch.Tensor:
"""
Low-level helper that assumes both inputs live on the same CUDA device and
are already contiguous and of dtype float16. Returns C on that device.
"""
if A_dev.dtype != torch.float16 or B_dev.dtype != torch.float16:
raise TypeError("Both A and B must be float16 tensors")
if A_dev.shape[1] != 2048:
raise ValueError(f"A must have second dimension 2048, got {A_dev.shape}")
if list(B_dev.shape) != [128, 2048]:
raise ValueError(f"B must have shape [128, 2048], got {B_dev.shape}")
# ----------- Tensor sizes & strides -------------------------------- #
M = A_dev.shape[0]
stride_am, stride_ak = A_dev.stride()
stride_bn, stride_bk = B_dev.stride()
C_dev = torch.empty((M, 128), dtype=torch.float16, device=A_dev.device)
stride_cm, stride_cn = C_dev.stride()
# ----------- Kernel configuration ---------------------------------- #
BLOCK_M = 64
BLOCK_N = 128 # covers the whole N dimension
BLOCK_K = 32
grid = (triton.cdiv(M, BLOCK_M),)
_gemm_n128_k2048_kernel[grid](
A_dev, B_dev, C_dev,
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,
)
return C_dev
# -----------------------------------------------------------------------------
# PUBLIC ENTRY POINT
# -----------------------------------------------------------------------------
def run(*args: Any, **kwargs: Dict[str, Any]) -> torch.Tensor:
"""
High-level helper replicating the reference API:
C = run(A, B) # positional
C = run(A=A_tensor, B=B) # keyword
Handles device management:
• Moves CPU tensors to GPU if necessary.
• Ensures both inputs are on the same device.
• Sends the result back to CPU if both inputs were on CPU.
"""
# ------------- Retrieve A and B arguments -------------------------- #
if len(args) >= 2:
A, B = args[:2]
else:
try:
A = kwargs["A"]
B = kwargs["B"]
except KeyError as exc:
raise ValueError("run expects tensors A and B either as positional "
"arguments or as keywords 'A' and 'B'") from exc
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required but not available")
# ------------- Decide target CUDA device --------------------------- #
if A.is_cuda and B.is_cuda:
target_device = A.device
if B.device != target_device:
raise RuntimeError("A and B must be on the same device")
elif A.is_cuda:
target_device = A.device
elif B.is_cuda:
target_device = B.device
else:
target_device = torch.device("cuda")
# ------------- Move inputs to GPU & make contiguous ----------------- #
A_dev = A.to(target_device, copy=False).contiguous()
B_dev = B.to(target_device, copy=False).contiguous()
# ------------- Launch the kernel ------------------------------------ #
C_dev = _launch_kernel(A_dev, B_dev)
# ------------- Move result back if inputs were on CPU --------------- #
if (not A.is_cuda) and (not B.is_cuda):
return C_dev.cpu()
return C_devscrolls · 180 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON