gpt-5-2025-08-07 / triton82d3cf
gpt-5-2025-08-07_triton_82d3cf · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 196 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-82d3cf?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
25 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 25 measurements ›Showing all 25 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:83de24e1d5e85c4974210fc44e8a0b8fcdeece984d871a7de4ea6afa9e48a188
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.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),mma
acc += tl.dot(a, tl.trans(b))num-warps = 8
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),stages = 5
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),Kernel source
main.py196 lines
import math
import torch
import triton
import triton.language as tl
# Autotuned GEMM kernel specialized for N=128, K=2048
# Computes: C[M, N] = A[M, K] @ B[N, K]^T
# A: [M, K] fp16, B: [N, K] fp16, C: [M, N] fp16
configs = [
# High-throughput default tile for Blackwell/Hopper-class GPUs
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),
# Smaller M tile for small/irregular M to improve occupancy
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=4, num_stages=5),
# Deeper K chunk for bandwidth-bound cases
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 128, "GROUP_M": 4}, num_warps=8, num_stages=4),
]
@triton.autotune(configs=configs, key=["M"])
@triton.jit
def gemm_n128_k2048_kernel(
A_ptr, B_ptr, C_ptr,
M, # runtime M
stride_am, stride_ak, # A strides
stride_bn, stride_bk, # B strides (N, K)
stride_cm, stride_cn, # C strides
K: tl.constexpr, # K compile-time constant (2048)
N: tl.constexpr, # N compile-time constant (128)
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr,
):
# Program ids for 2D launch grid
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# Offsets for M and N dimensions
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Accumulator in fp32 for numerical stability
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Base pointers for the first K-block
a_ptrs = A_ptr + (offs_m[:, None] * stride_am + tl.arange(0, BLOCK_K)[None, :] * stride_ak)
b_ptrs = B_ptr + (offs_n[:, None] * stride_bn + tl.arange(0, BLOCK_K)[None, :] * stride_bk)
# Masks for M/N boundaries, broadcast across K chunk
mask_m = offs_m[:, None] < M
mask_n = offs_n[:, None] < N
# K is guaranteed to be divisible by BLOCK_K for this problem (2048)
tl.static_assert(BLOCK_N == 128, "Kernel specialized for N tiles of 128.")
tl.static_assert((K % BLOCK_K) == 0, "K must be divisible by BLOCK_K.")
# Hint to compiler for better vectorization/coalescing
tl.max_contiguous(tl.arange(0, BLOCK_K), 64)
# Main K loop
for k0 in range(0, K, BLOCK_K):
a = tl.load(a_ptrs, mask=mask_m, other=0.0)
b = tl.load(b_ptrs, mask=mask_n, other=0.0)
# Compute partial matmul: (BM, BK) x (BK, BN)
acc += tl.dot(a, tl.trans(b))
# Advance pointers along K
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# Write back result (convert to fp16)
c_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
store_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, acc.to(tl.float16), mask=store_mask)
def _assert_and_prepare_inputs(A: torch.Tensor, B: torch.Tensor):
# Validate dtypes
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise TypeError("A and B must be torch.float16 (float16) tensors.")
# Validate ranks
if A.ndim != 2 or B.ndim != 2:
raise ValueError("A and B must be 2D tensors: A[M, K], B[N, K].")
M, K_a = A.shape
N, K_b = B.shape
if K_a != 2048 or K_b != 2048:
raise ValueError(f"K must be 2048. Got A.shape[1]={K_a}, B.shape[1]={K_b}.")
if N != 128:
raise ValueError(f"N must be 128. Got B.shape[0]={N}.")
return M, N, K_a
def _select_device_and_move(A: torch.Tensor, B: torch.Tensor):
# Determine computation device
a_dev = A.device
b_dev = B.device
cuda_available = torch.cuda.is_available()
# If any tensor is already on CUDA, use that device
if a_dev.type == "cuda" or b_dev.type == "cuda":
if not cuda_available:
raise RuntimeError("CUDA is not available but at least one input tensor is on CUDA.")
target_device = a_dev if a_dev.type == "cuda" else b_dev
A_dev = A.to(device=target_device, non_blocking=True)
B_dev = B.to(device=target_device, non_blocking=True)
return target_device, A_dev, B_dev
# Both tensors on CPU
if not cuda_available:
raise RuntimeError("CUDA is not available; Triton kernel requires a CUDA-capable GPU.")
target_device = torch.device("cuda", 0)
A_dev = A.to(device=target_device, non_blocking=True)
B_dev = B.to(device=target_device, non_blocking=True)
return target_device, A_dev, B_dev
def _launch_kernel(A_dev: torch.Tensor, B_dev: torch.Tensor, M: int, N: int, K: int):
# Ensure contiguity for optimal memory access
if not A_dev.is_contiguous():
A_dev = A_dev.contiguous()
if not B_dev.is_contiguous():
B_dev = B_dev.contiguous()
# Allocate output on device
C_dev = torch.empty((M, N), dtype=torch.float16, device=A_dev.device)
# Compute grid
def grid(meta):
return (
triton.cdiv(M, meta["BLOCK_M"]),
triton.cdiv(N, meta["BLOCK_N"]),
)
gemm_n128_k2048_kernel[grid](
A_dev, B_dev, C_dev,
M,
A_dev.stride(0), A_dev.stride(1),
B_dev.stride(0), B_dev.stride(1),
C_dev.stride(0), C_dev.stride(1),
K=K,
N=N,
)
return C_dev
def run(*args, **kwargs):
"""
Entry point: C = run(A, B)
- Computes C = A @ B.T for:
A: [M, 2048] float16
B: [128, 2048] float16
C: [M, 128] float16
- Optimized Triton kernel tuned for NVIDIA B200-class GPUs.
- Handles device management: moves inputs to GPU, runs kernel, and returns result on the original device of A if both inputs shared the same original device; otherwise returns on the compute device.
Args:
A (torch.Tensor): shape [M, 2048], dtype float16
B (torch.Tensor): shape [128, 2048], dtype float16
Returns:
torch.Tensor: shape [M, 128], dtype float16
"""
# Extract inputs from args/kwargs
if len(args) == 2:
A, B = args
else:
if "A" in kwargs and "B" in kwargs:
A = kwargs["A"]
B = kwargs["B"]
else:
raise ValueError("run expects two tensors: A and B, either as positional or keyword arguments.")
# Validate shapes/dtypes
M, N, K = _assert_and_prepare_inputs(A, B)
# Remember original devices
orig_dev_A = A.device
orig_dev_B = B.device
# Move to appropriate device
device, A_dev, B_dev = _select_device_and_move(A, B)
# Launch Triton kernel
C_dev = _launch_kernel(A_dev, B_dev, M, N, K)
# Move result back
# If both inputs were originally on the same device, return result to that device.
# Otherwise, return on the compute (CUDA) device.
if orig_dev_A == orig_dev_B:
target_out_dev = orig_dev_A
else:
target_out_dev = device
if target_out_dev.type == "cuda":
return C_dev.to(device=target_out_dev, non_blocking=True)
else:
return C_dev.cpu()scrolls · 196 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON