gemini-2.5-pro / tritonkmw3sz
gemini-2.5-pro_triton_kmw3sz · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 223 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-kmw3sz?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:46fb1691931ca4364810776666c3f1bbe3742b052babb7ae79b7efbc79723137
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.
autotune
@triton.autotune(mma
accumulator += tl.dot(a, b)tile-n = 128
BLOCK_SIZE_N=128, # N is fixed, so we set BLOCK_SIZE_N to the full dimension.Kernel source
main.py223 lines
import torch
import triton
import triton.language as tl
import math
# This Triton kernel is a fused GEMM operation C = A @ B.T
# A is of shape (M, K), B is of shape (N, K) -> C is of shape (M, N)
# It is optimized for the case where N=128 and K=2048.
# The M dimension is variable.
# Optimizations for B200 (Hopper-like architecture):
# - Use tl.dot() which maps to Tensor Core operations for fp16 inputs.
# - Accumulator is in fp32 for precision.
# - Autotuner explores different block sizes and pipeline stages.
# - BLOCK_SIZE_N is fixed to 128 to compute a full output row-block per thread block.
# - A 1D grid is used over the M-dimension, which is simple and effective.
# - Software pipelining is enabled via num_stages to hide memory latency.
@triton.autotune(
configs=[
# Basic configurations with varying block sizes and stages
triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_K': 64, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 32, 'num_stages': 5, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 64, 'num_stages': 3, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'num_stages': 2, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_K': 32, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_K': 64, 'num_stages': 3, 'num_warps': 8}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_K': 128, 'num_stages': 2, 'num_warps': 8}),
# Configurations with larger M-block for potentially large M inputs
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_K': 32, 'num_stages': 3, 'num_warps': 8}),
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_K': 64, 'num_stages': 2, 'num_warps': 8}),
],
key=['M'],
)
@triton.jit
def gemm_kernel(
# Pointers to matrices
A_ptr, B_ptr, C_ptr,
# Matrix dimensions
M, N, K,
# Strides for matrices
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
# Meta-parameters
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
):
"""
Triton kernel for GEMM: C = A @ B.T
A: [M, K]
B: [N, K]
C: [M, N]
"""
# -----------------------------------------------------------
# Map program ids to M-dimension
# This program will compute a BLOCK_SIZE_M x N tile of C
pid = tl.program_id(axis=0)
# -----------------------------------------------------------
# Create offsets for the M, N, and K dimensions.
# We will compute a BLOCK_SIZE_M x BLOCK_SIZE_N block of C.
offs_m = (pid * BLOCK_SIZE_M) + tl.arange(0, BLOCK_SIZE_M)
offs_n = tl.arange(0, BLOCK_SIZE_N)
offs_k = tl.arange(0, BLOCK_SIZE_K)
# -----------------------------------------------------------
# Initialise pointers to the first element of the A and B tiles.
# A is [M, K], B is [N, K].
# Pointer for A tile: [BLOCK_SIZE_M, BLOCK_SIZE_K]
A_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
# Pointer for B tile: [BLOCK_SIZE_K, BLOCK_SIZE_N]
# To compute A @ B.T, we need to effectively transpose the tile of B
# during the load. We do this by swapping the roles of N and K offsets
# in the pointer calculation.
B_ptrs = B_ptr + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)
# -----------------------------------------------------------
# Accumulator for the C tile, initialized to zeros.
# Using float32 for higher precision during accumulation.
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# -----------------------------------------------------------
# Loop over the K dimension by increments of BLOCK_SIZE_K.
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
# Load the tiles of A and B from global memory.
# Masking is applied to handle cases where M or K are not perfect multiples of block sizes.
mask_a = (offs_m[:, None] < M) & ((k * BLOCK_SIZE_K + offs_k[None, :]) < K)
a = tl.load(A_ptrs, mask=mask_a, other=0.0)
# Load a tile of B. Because of the pointer setup, this tile is
# effectively transposed, with shape [BLOCK_SIZE_K, BLOCK_SIZE_N].
# Masking is only needed for the K dimension.
mask_b = (k * BLOCK_SIZE_K + offs_k[:, None]) < K
b = tl.load(B_ptrs, mask=mask_b, other=0.0)
# Perform matrix multiplication.
# a: [BLOCK_SIZE_M, BLOCK_SIZE_K]
# b: [BLOCK_SIZE_K, BLOCK_SIZE_N]
# The result is [BLOCK_SIZE_M, BLOCK_SIZE_N], which is correct.
accumulator += tl.dot(a, 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 the output dtype (float16) and store the result.
c = accumulator.to(tl.float16)
# Create pointers for the C matrix and store the result.
offs_cm = (pid * BLOCK_SIZE_M) + tl.arange(0, BLOCK_SIZE_M)
offs_cn = tl.arange(0, BLOCK_SIZE_N)
C_ptrs = C_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
# Mask for writing to C to handle M not being a multiple of BLOCK_SIZE_M.
mask_c = (offs_cm[:, None] < M)
tl.store(C_ptrs, c, mask=mask_c)
def gemm_n128_k2048(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Wrapper function for the GEMM kernel.
Args:
A (torch.Tensor): A tensor of shape [M, 2048] and dtype float16.
B (torch.Tensor): A tensor of shape [128, 2048] and dtype float16.
Returns:
torch.Tensor: The result of A @ B.T, with shape [M, 128] and dtype float16.
"""
# --- Input validation ---
if not A.is_cuda or not B.is_cuda:
raise ValueError("Input tensors must be on a CUDA device.")
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise ValueError("Input tensors must have dtype torch.float16.")
M, K_A = A.shape
N, K_B = B.shape
if K_A != 2048 or K_B != 2048:
raise ValueError(f"K dimension must be 2048, but got {K_A} for A and {K_B} for B.")
if N != 128:
raise ValueError(f"N dimension for B must be 128, but got {N}.")
# --- Output tensor allocation ---
# The output C will have shape [M, N]
C = torch.empty((M, N), device=A.device, dtype=torch.float16)
# --- Grid setup ---
# The grid is 1D, with each program instance computing a block of M rows.
grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']),)
# --- Kernel launch ---
gemm_kernel[grid](
A, B, C,
M, N, 2048,
A.stride(0), A.stride(1),
B.stride(0), B.stride(1),
C.stride(0), C.stride(1),
BLOCK_SIZE_N=128, # N is fixed, so we set BLOCK_SIZE_N to the full dimension.
)
return C
def run(*args, **kwargs):
"""
Public entry point for the GEMM operation.
This function handles device management, argument parsing, and kernel execution.
It moves data to the GPU if necessary, runs the computation, and moves the
result back to the original device.
Usage:
run(A, B)
run(A=tensor_a, B=tensor_b)
"""
# --- Argument parsing ---
if 'A' in kwargs and 'B' in kwargs:
A = kwargs['A']
B = kwargs['B']
elif len(args) == 2:
A, B = args
else:
raise ValueError("run() expects two positional arguments (A, B) or two keyword arguments (A=..., B=...).")
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("Triton requires a CUDA-enabled GPU, but CUDA is not available.")
original_devices = {
'A': A.device,
'B': B.device,
}
# Determine the target CUDA device for computation.
# If any input is on CUDA, use that device. Otherwise, use the default CUDA device.
target_device = torch.device("cuda")
for tensor in [A, B]:
if tensor.is_cuda:
target_device = tensor.device
break
# Move all tensors to the target CUDA device.
try:
A_gpu = A.to(target_device)
B_gpu = B.to(target_device)
except Exception as e:
raise RuntimeError(f"Failed to move tensors to device {target_device}. Error: {e}")
# --- Execute the kernel ---
C_gpu = gemm_n128_k2048(A_gpu, B_gpu)
# --- Restore device of output ---
# The output tensor is moved back to the device of the first input tensor 'A'.
output_device = original_devices['A']
if C_gpu.device != output_device:
C = C_gpu.to(output_device)
else:
C = C_gpu
return Cscrolls · 223 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON