gpt-5-2025-08-07 / tritonffc694
gpt-5-2025-08-07_triton_ffc694 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 216 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-ffc694?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:27e8fd799035035fd7df4cdd5e8b660505d38f90519afe31140d65a795cb8e39
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, out_dtype=tl.float32)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.py216 lines
import math
import torch
import triton
import triton.language as tl
@triton.autotune(
configs=[
# Use power-of-two tile sizes to satisfy tl.arange power-of-two range requirement.
triton.Config({"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 64}, num_warps=8, num_stages=4),
triton.Config({"BLOCK_M": 128, "BLOCK_N": 512, "BLOCK_K": 64}, num_warps=8, num_stages=4),
triton.Config({"BLOCK_M": 64, "BLOCK_N": 256, "BLOCK_K": 128}, num_warps=4, num_stages=4),
triton.Config({"BLOCK_M": 256, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=4),
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=4),
triton.Config({"BLOCK_M": 256, "BLOCK_N": 256, "BLOCK_K": 64}, num_warps=8, num_stages=5),
],
key=["M"], # Autotune over M; N=5120, K=2048 are fixed
)
@triton.jit
def gemm_n5120_k2048_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[M, N] = A[M, K] @ B[N, K]^T
A: [M, K] row-major (stride_am, stride_ak)
B: [N, K] row-major (stride_bn, stride_bk) but we read B^T tiles
C: [M, N] row-major (stride_cm, stride_cn)
"""
# 2D launch grid over (M-tiles, N-tiles)
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
# Offsets for current 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)
# Help the compiler with alignment assumptions
tl.multiple_of(offs_m, 16)
tl.multiple_of(offs_n, 16)
tl.multiple_of(offs_k, 16)
# Initialize accumulation in FP32
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Base pointers for the first K-slice
a_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak) # [BM, BK]
# Load B as KxN by addressing B[n, k] -> B^T[k, n]
b_ptrs = B_ptr + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn) # [BK, BN]
# Iterate along K dimension
for k in range(0, K, BLOCK_K):
a_mask = (offs_m[:, None] < M) & (offs_k[None, :] + k < K)
b_mask = (offs_k[:, None] + k < K) & (offs_n[None, :] < N)
a = tl.load(a_ptrs, mask=a_mask, other=0.0)
b = tl.load(b_ptrs, mask=b_mask, other=0.0)
# Tensor Core accelerated: fp16 x fp16 -> fp32 accumulation
acc += tl.dot(a, b, out_dtype=tl.float32)
# Advance pointers to next K block
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# Write back results (convert to fp16)
c_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, acc.to(tl.float16), mask=c_mask)
def _assert_and_normalize_inputs(A: torch.Tensor, B: torch.Tensor):
if A is None or B is None:
raise ValueError("Expected tensors A and B, got None.")
if A.ndim != 2 or B.ndim != 2:
raise ValueError(f"Expected 2D tensors for A and B, got A.ndim={A.ndim}, B.ndim={B.ndim}")
M, K_a = A.shape
N_b, K_b = B.shape
if K_a != 2048:
raise ValueError(f"K dimension of A must be 2048, got {K_a}")
if K_b != 2048:
raise ValueError(f"K dimension of B (second dim) must be 2048, got {K_b}")
if N_b != 5120:
raise ValueError(f"N dimension of B (first dim) must be 5120, got {N_b}")
# Convert dtypes if needed
if A.dtype != torch.float16:
A = A.to(torch.float16)
if B.dtype != torch.float16:
B = B.to(torch.float16)
# Ensure contiguous layout (row-major) for efficient strided access
if not A.is_contiguous():
A = A.contiguous()
if not B.is_contiguous():
B = B.contiguous()
return A, B
def _call_triton_gemm(A: torch.Tensor, B: torch.Tensor, *, stream: torch.cuda.Stream | None = None):
"""
Internal: launch Triton kernel. The 'stream' argument is accepted for API
compatibility but not passed to Triton (Triton uses the current stream).
"""
# Shapes
M, K = A.shape
N = B.shape[0] # 5120 by contract
# Allocate output on same device as inputs (GPU)
C = torch.empty((M, N), dtype=torch.float16, device=A.device)
# Extract strides (in elements)
stride_am, stride_ak = A.stride()
stride_bn, stride_bk = B.stride()
stride_cm, stride_cn = C.stride()
# Grid: 2D grid over M-tiles and N-tiles
def grid(meta):
BM = meta["BLOCK_M"]
BN = meta["BLOCK_N"]
return (triton.cdiv(M, BM), triton.cdiv(N, BN))
# Launch kernel; do NOT pass 'stream' kwarg to Triton
gemm_n5120_k2048_kernel[grid](
A, B, C,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
)
return C
def run(*args, **kwargs):
"""
Entry point: C = run(A, B, stream=None)
Computes C = A @ B.T for:
- A: [M, 2048] float16
- B: [5120, 2048] float16
- C: [M, 5120] float16
Device management:
- If inputs are on CPU and CUDA is available, they are moved to GPU for the Triton kernel
- If any input is on CUDA but CUDA is not available, raises a clear error
- Result is moved back to the device of A (first input), preserving original device
- If CUDA is not available and both inputs are CPU tensors, falls back to torch.matmul on CPU
- Optional 'stream' (torch.cuda.Stream) sets the current stream for copies and compute
"""
# Unpack inputs from args/kwargs
if len(args) >= 2:
A, B = args[0], args[1]
else:
A = kwargs.get("A", None)
B = kwargs.get("B", None)
# Optional CUDA stream
stream = kwargs.get("stream", None)
# Validate shapes/dtypes and ensure contiguous layout
A, B = _assert_and_normalize_inputs(A, B)
a_dev = A.device
b_dev = B.device
cuda_available = torch.cuda.is_available()
# CPU-only path
if not cuda_available:
if A.is_cuda or B.is_cuda:
raise RuntimeError("CUDA tensor provided but CUDA is not available.")
return torch.matmul(A, B.T)
# CUDA available: ensure tensors on CUDA, respecting the provided stream
if stream is not None:
if not isinstance(stream, torch.cuda.Stream):
raise TypeError("stream must be a torch.cuda.Stream or None")
with torch.cuda.stream(stream):
A_cuda = A if A.is_cuda else A.cuda(non_blocking=True)
B_cuda = B if B.is_cuda else B.cuda(non_blocking=True)
C_cuda = _call_triton_gemm(A_cuda, B_cuda, stream=stream) # stream is ignored by Triton
# Move result back to match A's original device (if needed)
if a_dev.type == "cuda":
C_out = C_cuda
else:
C_out = C_cuda.to(a_dev, non_blocking=True)
else:
# Default stream
A_cuda = A if A.is_cuda else A.cuda(non_blocking=True)
B_cuda = B if B.is_cuda else B.cuda(non_blocking=True)
C_cuda = _call_triton_gemm(A_cuda, B_cuda, stream=None)
if a_dev.type == "cuda":
C_out = C_cuda
else:
C_out = C_cuda.to(a_dev, non_blocking=True)
return C_out
if __name__ == "__main__":
# Simple sanity check
torch.manual_seed(0)
M = 512 # example M
A = torch.randn((M, 2048), dtype=torch.float16)
B = torch.randn((5120, 2048), dtype=torch.float16)
C = run(A, B)
ref = torch.matmul(A, B.T)
max_diff = (C.cpu().float() - ref.float()).abs().max().item()
print("Max abs diff:", max_diff)scrolls · 216 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON