gemini-2.5-pro / triton5iu7uf
gemini-2.5-pro_triton_5iu7uf · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 165 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-5iu7uf?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:5df2577f34e44188492725e911f0e23756bb00a8169f059c9f232208b13c1bd4
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
main.py165 lines
import torch
import triton
import triton.language as tl
import math
@triton.autotune(
configs=[
# Basic configurations
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'num_warps': 4, 'num_stages': 2}, num_ctas=1),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'num_warps': 8, 'num_stages': 2}, num_ctas=1),
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'num_warps': 8, 'num_stages': 2}, num_ctas=1),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 64, 'num_warps': 4, 'num_stages': 3}, num_ctas=1),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'num_warps': 4, 'num_stages': 3}, num_ctas=1),
# Configurations potentially good for B200 with larger compute/memory resources
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'num_warps': 8, 'num_stages': 3}, num_ctas=1),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'num_warps': 8, 'num_stages': 3}, num_ctas=1),
],
key=['M'],
)
@triton.jit
def gemm_kernel(
A, B, C,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
):
"""
Computes C = A @ B.T
A is of shape (M, K)
B is of shape (N, K)
C is of shape (M, N)
This kernel is optimized for a matrix multiplication where the second matrix (B)
is transposed. Both A and B are expected to be row-major.
The kernel is structured to perform coalesced loads from both A and B.
The transpose operation is handled by `tl.trans` on the register-loaded tile of B
before the `tl.dot` operation. This approach relies on the compiler to efficiently
schedule the transpose and dot instructions and is effective on modern GPUs with
large caches like B200.
"""
# -----------------------------------------------------------
# Map program ids to M and N dimensions.
pid = tl.program_id(axis=0)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
# ----------------------------------------------------------
# Create pointers for the first blocks of A and B.
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
offs_k = tl.arange(0, BLOCK_SIZE_K)
a_ptrs = A + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
b_ptrs = B + (offs_n[:, None] * stride_bn + offs_k[None, :] * stride_bk)
# -----------------------------------------------------------
# Initialize accumulator.
# The accumulator is in float32 to maintain precision.
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# -----------------------------------------------------------
# Loop over the K dimension.
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
# Load the next block of A and B.
# Masking is applied to handle the case where M is not a multiple of BLOCK_SIZE_M.
# Since N and K are constants and our block sizes divide them, masks for N and K
# are not strictly necessary but are kept for generality. The compiler will optimize
# them out if possible.
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K), other=0.0)
b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_k[None, :] < K), other=0.0)
# The operation is C = A @ B.T, which translates to C[m,n] = sum_k A[m,k] * B[n,k].
# Our loaded tile `a` is (BLOCK_M, BLOCK_K) and `b` is (BLOCK_N, BLOCK_K).
# We need to compute dot(a, b.T). tl.trans(b) makes it (BLOCK_K, BLOCK_N).
accumulator += tl.dot(a, tl.trans(b))
# Advance pointers to the next K block.
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += BLOCK_SIZE_K * stride_bk
# -----------------------------------------------------------
# Cast accumulator to output dtype and write back to C.
C_out = accumulator.to(tl.float16)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
c_ptrs = C + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, C_out, mask=c_mask)
def run(*args, **kwargs):
"""
Wrapper function for the GEMM kernel, providing a user-friendly interface
and handling all device management.
Args:
A (torch.Tensor): The first input tensor of shape [M, K].
B (torch.Tensor): The second input tensor of shape [N, K].
Can be passed as positional or keyword arguments.
Returns:
torch.Tensor: The output tensor C of shape [M, N], on the same device as the input A.
"""
# --- Argument parsing ---
if len(args) == 2:
A, B = args
elif 'A' in kwargs and 'B' in kwargs:
A = kwargs['A']
B = kwargs['B']
else:
raise ValueError("Inputs 'A' and 'B' must be provided either as positional or keyword arguments.")
# --- Shape and DType validation ---
assert A.dtype == torch.float16, f"Input A must be float16, but got {A.dtype}"
assert B.dtype == torch.float16, f"Input B must be float16, but got {B.dtype}"
M, K_A = A.shape
N, K_B = B.shape
# Constants from the spec
spec_N, spec_K = 5120, 2048
assert K_A == spec_K, f"A.shape[1] must be {spec_K}, but got {K_A}"
assert K_B == spec_K, f"B.shape[1] must be {spec_K}, but got {K_B}"
assert N == spec_N, f"B.shape[0] must be {spec_N}, but got {N}"
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
original_device = A.device
device = torch.device("cuda")
# Move tensors to GPU if they are not already there
A_gpu = A.to(device, non_blocking=True) if A.device != device else A
B_gpu = B.to(device, non_blocking=True) if B.device != device else B
# Ensure inputs are contiguous for optimal memory access
A_gpu = A_gpu.contiguous()
B_gpu = B_gpu.contiguous()
# --- Output Tensor Allocation ---
C = torch.empty((M, N), device=device, dtype=torch.float16)
# --- Grid Definition ---
# We use a 1D grid to simplify launching and autotuning, especially for the dynamic M dimension.
# The kernel then internally maps the 1D program ID to 2D (M, N) block coordinates.
grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']),)
# --- Kernel Launch ---
gemm_kernel[grid](
A_gpu, B_gpu, C,
M, N, spec_K,
A_gpu.stride(0), A_gpu.stride(1),
B_gpu.stride(0), B_gpu.stride(1),
C.stride(0), C.stride(1),
)
# --- Result Handling ---
# Move the result back to the original device of the input tensor
return C.to(original_device, non_blocking=True)scrolls · 165 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON