gpt-5-2025-08-07 / triton9b01eb
gpt-5-2025-08-07_triton_9b01eb · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 138 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-9b01eb?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
29 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 29 measurements ›Showing all 29 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:e233d827702d82b001ae4ce0cc006a334d047bf2a3a6bcedd79b168eb75cdc6c
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=4),stages = 4
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=4),Kernel source
main.py138 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=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 128}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=4, num_stages=4),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=4, num_stages=3),
],
key=['M'],
)
@triton.jit
def gemm_n2048_k4096_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,
):
# Program ids for 2D launch grid
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
# Offsets for the current block
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 the first K-slice of A and B for this tile
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
# Accumulator in FP32 for improved accuracy
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K dimension
k = 0
while k < K:
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (k + offs_k[None, :] < K), other=0.0)
b = tl.load(b_ptrs, mask=(offs_n[None, :] < N) & (k + offs_k[:, None] < K), other=0.0)
acc += tl.dot(a, b)
k += BLOCK_K
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# Write back to C (cast to FP16)
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 _validate_inputs(A: torch.Tensor, B: torch.Tensor):
if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
raise TypeError("Inputs A and B must be torch.Tensor instances.")
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise TypeError(f"Inputs must be torch.float16. Got A.dtype={A.dtype}, B.dtype={B.dtype}.")
if A.ndim != 2 or B.ndim != 2:
raise ValueError(f"Inputs must be 2D tensors. Got A.ndim={A.ndim}, B.ndim={B.ndim}.")
M, K_a = A.shape
N_b, K_b = B.shape
if K_a != 4096:
raise ValueError(f"A must have shape [M, 4096]. Got {A.shape}.")
if N_b != 2048 or K_b != 4096:
raise ValueError(f"B must have shape [2048, 4096]. Got {B.shape}.")
return M, 2048, 4096
def _compute_device_for_inputs(A: torch.Tensor, B: torch.Tensor):
a_dev = A.device
b_dev = B.device
cuda_available = torch.cuda.is_available()
# If any input is CUDA, ensure CUDA is available and devices match
if a_dev.type == 'cuda' or b_dev.type == 'cuda':
if not cuda_available:
raise RuntimeError("CUDA tensor provided but CUDA is not available.")
if a_dev.type == 'cuda' and b_dev.type == 'cuda' and a_dev != b_dev:
raise ValueError("A and B must be on the same CUDA device.")
return a_dev if a_dev.type == 'cuda' else b_dev
# Both on CPU
if not cuda_available:
raise RuntimeError("CUDA is required for this Triton kernel, but no CUDA device is available.")
# Use current CUDA device
idx = torch.cuda.current_device()
return torch.device(f"cuda:{idx}")
def run(*args, **kwargs):
# Extract A and B from args/kwargs
if len(args) >= 2:
A, B = args[0], args[1]
else:
if 'A' not in kwargs or 'B' not in kwargs:
raise ValueError("run requires tensors A and B either as positional or keyword arguments.")
A, B = kwargs['A'], kwargs['B']
M, N, K = _validate_inputs(A, B)
compute_device = _compute_device_for_inputs(A, B)
# Track original devices to restore output
orig_a_dev = A.device
orig_b_dev = B.device
return_to_cpu = (orig_a_dev.type != 'cuda') and (orig_b_dev.type != 'cuda')
# Move inputs to compute device and ensure contiguous layout for best performance
with torch.cuda.device(compute_device.index if compute_device.index is not None else 0):
A_dev = A.to(device=compute_device, non_blocking=True)
B_dev = B.to(device=compute_device, non_blocking=True)
# Contiguous for coalesced memory accesses
if not A_dev.is_contiguous():
A_dev = A_dev.contiguous()
if not B_dev.is_contiguous():
B_dev = B_dev.contiguous()
# Allocate output
C_dev = torch.empty((M, N), device=compute_device, dtype=torch.float16)
# Launch kernel
grid = lambda META: (triton.cdiv(M, META['BLOCK_M']), triton.cdiv(N, META['BLOCK_N']))
gemm_n2048_k4096_kernel[grid](
A_dev, B_dev, C_dev,
M, N, K,
A_dev.stride(0), A_dev.stride(1),
B_dev.stride(0), B_dev.stride(1),
C_dev.stride(0), C_dev.stride(1),
)
# Restore output to original device(s)
if return_to_cpu:
return C_dev.cpu()
else:
# Keep on GPU device where inputs lived
return C_devscrolls · 138 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON