gpt-5-2025-08-07 / triton8c14a2
gpt-5-2025-08-07_triton_8c14a2 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 178 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-8c14a2?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:8b38711a03b00ced0e765466e7c49b64224b5fa0c92786a0d080e0364295062b
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': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),stages = 4
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),Kernel source
main.py178 lines
import torch
import triton
import triton.language as tl
# Autotuned GEMM kernel for:
# A: [M, K] fp16
# B: [N, K] fp16
# C: [M, N] fp16
# Computes: C = A @ B.T
# Optimized for N=256, K=7168; M is variable.
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 128, 'GROUP_M': 8}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 256, 'BLOCK_K': 128, 'GROUP_M': 8}, num_warps=4, num_stages=5),
],
key=['M'], # N and K are constant for this op
)
@triton.jit
def _gemm_n256_k7168_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,
GROUP_M: tl.constexpr,
):
# Program IDs with swizzled 1D launch to improve L2 hit-rate across M
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_in_group = GROUP_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_M
group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_M)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
# Offsets for this tile
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)
# Pointers to tiles of A and B; note B is [N, K], we load [BK, BN] and use dot(A[BM,BK], B[BK,BN])
A_tile_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
B_tile_ptrs = B_ptr + (offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk)
# Accumulator in FP32
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Masks for bounds
m_mask = offs_m < M
n_mask = offs_n < N
# K loop
k_iter = 0
while k_iter < K:
k_mask = (k_iter + offs_k) < K
a = tl.load(A_tile_ptrs, mask=m_mask[:, None] & k_mask[None, :], other=0.0)
b = tl.load(B_tile_ptrs, mask=k_mask[:, None] & n_mask[None, :], other=0.0)
acc += tl.dot(a, b)
# Advance pointers along K
A_tile_ptrs += BLOCK_K * stride_ak
B_tile_ptrs += BLOCK_K * stride_bk
k_iter += BLOCK_K
# Write back in FP16
C_tile_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
tl.store(C_tile_ptrs, acc.to(tl.float16), mask=m_mask[:, None] & n_mask[None, :])
def _validate_inputs(A: torch.Tensor, B: torch.Tensor):
if A.ndim != 2 or B.ndim != 2:
raise ValueError(f"Expected 2D tensors, got A.ndim={A.ndim}, B.ndim={B.ndim}")
M, K_a = A.shape
N_b, K_b = B.shape
if N_b != 256:
raise ValueError(f"N must be 256; got B.shape[0]={N_b}")
if K_a != 7168 or K_b != 7168:
raise ValueError(f"K must be 7168; got A.shape[1]={K_a}, B.shape[1]={K_b}")
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise TypeError(f"Expected A and B to be torch.float16; got {A.dtype} and {B.dtype}")
return M, 256, 7168
def _select_run_device(A: torch.Tensor, B: torch.Tensor):
# If CUDA not available
if not torch.cuda.is_available():
if A.is_cuda or B.is_cuda:
raise RuntimeError("CUDA is not available, but at least one input tensor is on CUDA.")
return None # CPU-only scenario
# CUDA available
if A.is_cuda and B.is_cuda:
if A.device != B.device:
raise ValueError(f"A and B must be on the same CUDA device; got {A.device} and {B.device}")
return A.device
# If only one is CUDA, use that device; if both CPU, use current CUDA device
if A.is_cuda:
return A.device
if B.is_cuda:
return B.device
return torch.device('cuda', torch.cuda.current_device())
def _launch_triton(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
M, N, K = _validate_inputs(A, B)
run_device = _select_run_device(A, B)
# CPU-only fallback if CUDA not available
if run_device is None:
# Keep CPU semantics
return torch.matmul(A, B.t())
# Move inputs to the chosen CUDA device (non-blocking when possible)
A_dev = A.to(device=run_device, non_blocking=True)
B_dev = B.to(device=run_device, non_blocking=True)
# Allocate output on device
C_dev = torch.empty((M, N), dtype=torch.float16, device=run_device)
# Strides in elements
stride_am, stride_ak = A_dev.stride()
stride_bn, stride_bk = B_dev.stride()
stride_cm, stride_cn = C_dev.stride()
# Grid: 1D with grouping across M to improve cache locality on B200
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
)
# Launch kernel
_gemm_n256_k7168_kernel[grid](
A_dev, B_dev, C_dev,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
)
# Restore output device semantics:
# - If both inputs were CPU originally: return CPU tensor
# - Else if any input was CUDA originally: return on that CUDA device (first CUDA in A,B order)
if (not A.is_cuda) and (not B.is_cuda):
return C_dev.to(device=A.device, non_blocking=True)
target_out = A.device if A.is_cuda else (B.device if B.is_cuda else run_device)
if target_out != run_device:
return C_dev.to(device=target_out, non_blocking=True)
return C_dev
def run(*args, **kwargs):
"""
Entry point matching the reference signature.
Usage:
- run(A, B)
- run(A=A, B=B)
This will:
- Move CPU tensors to GPU if CUDA is available
- Validate shapes/dtypes (M,256,7168; fp16)
- Launch a Triton-optimized GEMM for B200
- Return result on the original device (CPU if inputs were CPU-only; otherwise on the input CUDA device)
"""
if len(args) >= 2:
A, B = args[0], args[1]
else:
A = kwargs.get('A', None)
B = kwargs.get('B', None)
if A is None or B is None:
raise ValueError("run requires tensors A and B, either as positional args or kwargs.")
return _launch_triton(A, B)scrolls · 178 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON