gpt-5 / triton998d17
gpt-5_triton_998d17 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 169 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-998d17?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:3bde1487665689bc3e962fc344c2df1d862c5031ff03314d3d109da0ca27199a
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
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
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),stages = 5
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),Kernel source
main.py169 lines
import torch
import triton
import triton.language as tl
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=5),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=4, num_stages=5),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=4, num_stages=5),
triton.Config({'BLOCK_M': 256, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 128}, num_warps=8, num_stages=4),
],
key=['M'],
)
@triton.jit
def _gemm_n_28672_k_4096_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,
):
tl.static_assert(BLOCK_K % 16 == 0)
tl.static_assert(BLOCK_M % 16 == 0)
tl.static_assert(BLOCK_N % 16 == 0)
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = 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
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k = 0
while k < K:
a = tl.load(
a_ptrs,
mask=(offs_m[:, None] < M) & (offs_k[None, :] + k < K),
other=0.0,
)
b = tl.load(
b_ptrs,
mask=(offs_k[:, None] + k < K) & (offs_n[None, :] < N),
other=0.0,
)
acc += tl.dot(a, b)
k += BLOCK_K
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
c = acc.to(tl.float16)
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(
c_ptrs,
c,
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
)
def run(*args, **kwargs):
A = None
B = None
if len(args) >= 1:
A = args[0]
if len(args) >= 2:
B = args[1]
if 'A' in kwargs:
A = kwargs['A']
if 'B' in kwargs:
B = kwargs['B']
if A is None or B is None:
raise ValueError("run(A, B): both A and B must be provided")
if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
raise TypeError("A and B must be torch.Tensor")
# Validate dtypes and shapes
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise TypeError("A and B must be float16 tensors")
if A.ndim != 2 or B.ndim != 2:
raise ValueError("A and B must be 2D tensors")
M, K_a = A.shape
N_b, K_b = B.shape
N_SPEC = 28672
K_SPEC = 4096
if K_a != K_SPEC or K_b != K_SPEC:
raise ValueError(f"K dimension must be {K_SPEC}; got A.shape[1]={K_a}, B.shape[1]={K_b}")
if N_b != N_SPEC:
raise ValueError(f"B.shape[0] (N) must be {N_SPEC}; got {N_b}")
# Device management
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available; Triton kernel requires a CUDA-capable device")
# Choose compute device
compute_device = None
if A.is_cuda:
compute_device = A.device
if B.is_cuda:
# If both CUDA and different devices, prefer A's; else use B's
compute_device = A.device if A.is_cuda else B.device
if A.is_cuda and A.device != B.device:
# Move B to A's device to compute
pass
if compute_device is None:
compute_device = torch.device('cuda')
# Move inputs to compute_device if needed
if not A.is_cuda or A.device != compute_device:
A_dev = A.to(device=compute_device, dtype=torch.float16, non_blocking=True)
else:
A_dev = A
if not B.is_cuda or B.device != compute_device:
B_dev = B.to(device=compute_device, dtype=torch.float16, non_blocking=True)
else:
B_dev = B
# Prepare output on compute_device
C_dev = torch.empty((M, N_SPEC), device=compute_device, dtype=torch.float16)
# Strides (in elements)
stride_am = A_dev.stride(0)
stride_ak = A_dev.stride(1)
stride_bn = B_dev.stride(0)
stride_bk = B_dev.stride(1)
stride_cm = C_dev.stride(0)
stride_cn = C_dev.stride(1)
# Grid
def grid(meta):
return (
triton.cdiv(M, meta['BLOCK_M']),
triton.cdiv(N_SPEC, meta['BLOCK_N']),
)
_gemm_n_28672_k_4096_kernel[grid](
A_dev, B_dev, C_dev,
M, N_SPEC, K_SPEC,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
)
# Move result back to original device if both inputs were originally on CPU
# If any input was originally CUDA, return on that CUDA device (A's CUDA device takes precedence)
if (not A.is_cuda) and (not B.is_cuda):
return C_dev.cpu()
else:
# If A was originally CUDA and not on compute_device, move to A's original device?
# Requirement: preserve original tensor devices and restore them for outputs.
# Use A's original CUDA device if it was CUDA; else use B's original CUDA device.
target_device = A.device if A.is_cuda else (B.device if B.is_cuda else compute_device)
if C_dev.device != target_device:
return C_dev.to(target_device, non_blocking=True)
return C_devscrolls · 169 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON