gpt-o3 / tritonef0d11
gpt-o3_triton_ef0d11 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 118 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-ef0d11?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
Show all 43 measurements ›Showing all 43 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:5d846b9f7c2d09afe4cefb4461449998b40106ce852f5d4cae2d736dd2cb3728
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) # fp32 accumulationnum-warps = 8
num_warps=8,stages = 4
num_stages=4tile-k = 32
def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):tile-m = 128
def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):tile-n = 128
def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):Kernel source
main.py118 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def _gemm_4096_kernel(
A_ptr, B_ptr, C_ptr,
M, # total rows of A / C
stride_am, stride_ak, # strides for A
stride_bk, stride_bn, # strides for B (transposed view K x N)
stride_cm, stride_cn, # strides for C
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr
):
"""
Kernel computes C = A @ B, where
A: [M, K] (row-major)
B: [K, N] (row-major, supplied as B.T from host for contiguous K dimension)
K = 4096, N = 4096 (constants)
"""
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]
# Pointer to C
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
# Accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K axis
for k0 in range(0, 4096, BLOCK_K):
offs_k = k0 + tl.arange(0, 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_k[:, None] * stride_bk + offs_n[None, :] * stride_bn # (BLOCK_K, BLOCK_N)
a = tl.load(a_ptrs, mask=offs_m[:, None] < M, other=0.0) # mask for last incomplete block
b = tl.load(b_ptrs) # K and N are constants, no mask
acc += tl.dot(a, b) # fp32 accumulation
# Write back
c = acc.to(tl.float16)
tl.store(c_ptrs, c, mask=offs_m[:, None] < M)
def _launch_kernel(A, B_T, C, BLOCK_M=128, BLOCK_N=128, BLOCK_K=32):
# Grid dimensions
M = A.shape[0]
grid_m = triton.cdiv(M, BLOCK_M)
grid_n = 32 # 4096 / 128
_gemm_4096_kernel[(grid_m, grid_n)](
A, B_T, C,
M,
A.stride(0), A.stride(1),
B_T.stride(0), B_T.stride(1),
C.stride(0), C.stride(1),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
num_warps=8,
num_stages=4
)
def run(A: torch.Tensor, B: torch.Tensor):
"""
Entry point that matches reference semantics:
C = A @ B.T
Handles device placement transparently.
"""
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernel")
# Preserve original devices
device_a = A.device
device_b = B.device
# Move to GPU if necessary
A_gpu = A.cuda() if not A.is_cuda else A
B_gpu = B.cuda() if not B.is_cuda else B
# Materialize B^T as contiguous KxN tensor
B_T = B_gpu.t().contiguous()
# Allocate output tensor on GPU
M = A_gpu.shape[0]
C_gpu = torch.empty((M, 4096), device=A_gpu.device, dtype=torch.float16)
# Launch Triton kernel
_launch_kernel(A_gpu, B_T, C_gpu)
# Move result back to the device of A (arbitrary choice if A & B differ)
C_out = C_gpu.to(device_a)
return C_out
# If this file is executed directly, run a quick correctness test
if __name__ == "__main__":
torch.manual_seed(0)
M_test = 512
A_test = torch.randn((M_test, 4096), dtype=torch.float16)
B_test = torch.randn((4096, 4096), dtype=torch.float16)
C_ref = torch.matmul(A_test, B_test.t())
C_triton = run(A_test, B_test)
max_err = (C_ref - C_triton).abs().max()
print("Max error:", max_err.item())scrolls · 118 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON