gpt-o3_triton_189201
gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 162 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-189201?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
17 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 17 measurements ›Showing all 17 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:b6694ec82d45fe33e3d4bdd9bbb95cacabc84dfb6bcc5db38fc2de95286dea1c
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) # (BM, BN)num-warps = 8
num_warps=8, num_stages=4,stages = 4
num_warps=8, num_stages=4,tile-k = 128
BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 32Kernel source
main.py162 lines
import math
import torch
import triton
import triton.language as tl
# -----------------------------------------------------------------------------
# Problem constants (host-side)
# -----------------------------------------------------------------------------
K_TOTAL: int = 7_168 # fixed K dimension
N_CONST: int = 256 # fixed N dimension
DTYPE = torch.float16 # all tensors are fp16
# -----------------------------------------------------------------------------
# Triton kernel
# -----------------------------------------------------------------------------
@triton.jit
def _gemm_kernel(
A_ptr, B_ptr, C_ptr, # pointers to tensors
M, # run-time length of the M axis
stride_am, stride_ak, # strides for A
stride_bk, stride_bn, # strides for B (row-major [N,K])
stride_cm, stride_cn, # strides for C
BLOCK_M: tl.constexpr, # tile sizes (compile-time)
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""
Compute C = A @ B.T
A : [M, K_TOTAL]
B : [N_CONST, K_TOTAL] (row-major / torch-contiguous)
C : [M, N_CONST]
The kernel spans the whole N axis (256) inside a single program instance
and tiles the M axis with BLOCK_M.
"""
# ------------------------------------------------------------------
# Program-ID → position of the tile on the M axis
# ------------------------------------------------------------------
pid_m = tl.program_id(0)
m_start = pid_m * BLOCK_M
# ------------------------------------------------------------------
# Indices within the current tile
# ------------------------------------------------------------------
offs_m = m_start + tl.arange(0, BLOCK_M) # [BM]
offs_n = tl.arange(0, BLOCK_N) # [BN] (0 … 255)
offs_k = tl.arange(0, BLOCK_K) # [BK]
# Accumulator in fp32 for improved precision
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# ------------------------------------------------------------------
# Iterate over the K axis (compile-time constant length 7168)
# ------------------------------------------------------------------
for k_off in tl.static_range(0, 7_168, BLOCK_K):
k_idx = k_off + offs_k # [BK]
# Pointers for the current K-slice
a_ptrs = A_ptr + (offs_m[:, None] * stride_am +
k_idx[None, :] * stride_ak) # (BM, BK)
b_ptrs = B_ptr + (k_idx[:, None] * stride_bk +
offs_n[None, :] * stride_bn) # (BK, BN)
# Guard against incomplete last M-tile
mask_a = offs_m[:, None] < M
a = tl.load(a_ptrs, mask=mask_a, other=0.0) # (BM, BK)
b = tl.load(b_ptrs) # (BK, BN)
acc += tl.dot(a, b) # (BM, BN)
# ------------------------------------------------------------------
# Write results back to global memory
# ------------------------------------------------------------------
c = acc.to(tl.float16) # fp32 → fp16
c_ptrs = C_ptr + (offs_m[:, None] * stride_cm +
offs_n[None, :] * stride_cn)
tl.store(c_ptrs, c, mask=offs_m[:, None] < M)
# -----------------------------------------------------------------------------
# Helper that launches the Triton kernel
# -----------------------------------------------------------------------------
def _launch_gemm(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for Triton kernels but is not available.")
# Move to GPU if necessary
A_gpu = A.cuda() if A.device.type != "cuda" else A
B_gpu = B.cuda() if B.device.type != "cuda" else B
# Basic sanity checks
if A_gpu.dtype != DTYPE or B_gpu.dtype != DTYPE:
raise TypeError("A and B must be fp16 tensors.")
if A_gpu.shape[1] != K_TOTAL:
raise ValueError(f"A must have shape (M, {K_TOTAL})")
if B_gpu.shape != (N_CONST, K_TOTAL):
raise ValueError(f"B must have shape ({N_CONST}, {K_TOTAL})")
# Ensure contiguous layout for coalesced accesses
A_gpu = A_gpu.contiguous()
B_gpu = B_gpu.contiguous()
M = A_gpu.shape[0]
C_gpu = torch.empty((M, N_CONST), dtype=DTYPE, device=A_gpu.device)
# Kernel configuration
BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 32
grid = (triton.cdiv(M, BLOCK_M),) # 1-D grid along M
_gemm_kernel[grid](
A_gpu, B_gpu, C_gpu,
M,
A_gpu.stride(0), A_gpu.stride(1),
B_gpu.stride(1), B_gpu.stride(0),
C_gpu.stride(0), C_gpu.stride(1),
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
num_warps=8, num_stages=4,
)
return C_gpu
# -----------------------------------------------------------------------------
# Public entry point (mirrors the reference interface)
# -----------------------------------------------------------------------------
def run(*args, **kwargs):
"""
Usage:
C = run(A, B)
C = run(A=<tensor>, B=<tensor>)
The result tensor is returned on the same device as A.
"""
# ------------------------------------------------------ #
# Parse arguments
# ------------------------------------------------------ #
if len(args) + len(kwargs) != 2:
raise ValueError("run expects exactly two tensor arguments: A and B")
if len(args) == 2:
A, B = args
else:
A = kwargs.pop("A", None)
B = kwargs.pop("B", None)
if A is None or B is None:
raise ValueError("Both A and B must be provided.")
if kwargs:
raise ValueError(f"Unexpected keyword arguments: {tuple(kwargs.keys())}")
# ------------------------------------------------------ #
# Execute kernel
# ------------------------------------------------------ #
C_gpu = _launch_gemm(A, B)
# Return result on caller's original device
return C_gpu.to(A.device)
__all__ = ["run"]scrolls · 162 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON