gpt-o3_triton_a47205
gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 143 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-a47205?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:2347747be2b550c4de9eba44a19d87ae08554e8d4548b1c73a15f3e2cf795b71
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.
autotune
@triton.autotune(mma
acc += tl.dot(a, b)num-warps = 8
num_warps=8, num_stages=3),stages = 3
num_warps=8, num_stages=3),Kernel source
main.py143 lines
import math
import torch
import triton
import triton.language as tl
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64},
num_warps=8, num_stages=3),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 64},
num_warps=4, num_stages=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 64},
num_warps=4, num_stages=4),
],
key=['M'],
)
@triton.jit
def _gemm_kernel(A_ptr, B_ptr, C_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""
Compute C = A @ B.T
A: [M, K] row-major (stride_am, stride_ak)
B: [N, K] row-major (stride_bn, stride_bk) – will be accessed transposed
C: [M, N] row-major (stride_cm, stride_cn)
All tensors are float16, accumulation in float32.
"""
pid = tl.program_id(axis=0)
grid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid // grid_n
pid_n = pid % grid_n
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Initialize accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K dimension (K == 14336, divisible by BLOCK_K == 64)
for k0 in tl.static_range(0, 14336, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
b_ptrs = B_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk
mask_a = (offs_m[:, None] < M) & (offs_k[None, :] < K)
mask_b = (offs_n[None, :] < N) & (offs_k[:, None] < K)
a = tl.load(a_ptrs, mask=mask_a, other=0.).to(tl.float16)
b = tl.load(b_ptrs, mask=mask_b, other=0.).to(tl.float16)
acc += tl.dot(a, b)
# Write back result
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
mask_c = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, acc.to(tl.float16), mask=mask_c)
def _launch_kernel(A_gpu: torch.Tensor, B_gpu: torch.Tensor) -> torch.Tensor:
M, K = A_gpu.shape
N = B_gpu.shape[0] # 4096
C_gpu = torch.empty((M, N), device=A_gpu.device, dtype=torch.float16)
stride_am, stride_ak = A_gpu.stride()
stride_bn, stride_bk = B_gpu.stride()
stride_cm, stride_cn = C_gpu.stride()
def grid(meta):
return (
triton.cdiv(M, meta['BLOCK_M']) *
triton.cdiv(N, meta['BLOCK_N']),
)
_gemm_kernel[grid](
A_gpu, B_gpu, C_gpu,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
)
return C_gpu
def run(A: torch.Tensor, B: torch.Tensor):
"""
Entry point. Computes C = A @ B.T using a Triton kernel optimized for NVIDIA B200 GPUs.
Parameters
----------
A : torch.Tensor
Input tensor of shape [M, 14336] (float16)
B : torch.Tensor
Input tensor of shape [4096, 14336] (float16)
Returns
-------
torch.Tensor
Result tensor of shape [M, 4096] (float16) on the same device type as inputs.
"""
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required for Triton kernel execution.")
# Preserve original devices
orig_device_A = A.device
orig_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
# Shape validation
if A_gpu.dtype != torch.float16 or B_gpu.dtype != torch.float16:
raise TypeError("Input tensors must be float16.")
if A_gpu.shape[1] != 14336 or B_gpu.shape != (4096, 14336):
raise ValueError(
"Expected shapes: A [M, 14336], B [4096, 14336]; got "
f"A {tuple(A_gpu.shape)}, B {tuple(B_gpu.shape)}"
)
# Launch Triton kernel
C_gpu = _launch_kernel(A_gpu, B_gpu)
# Sync before moving results
torch.cuda.synchronize()
# Restore to original device context
if not A.is_cuda and not B.is_cuda:
return C_gpu.cpu()
return C_gpu
__all__ = ["run"]scrolls · 143 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON