gemini-2.5-pro / tritonnekk4o
gemini-2.5-pro_triton_nekk4o · gemini-2.5-pro · 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-gemini-2-5-pro-triton-nekk4o?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:cfdc62bc0b2ab6a656f937835ddf9debeecfeb72049ec5db049a4ae9f3799d0c
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.py196 lines
import torch
import triton
import triton.language as tl
import math
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 2, 'num_warps': 4}),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 32, 'num_stages': 2, 'num_warps': 4}),
# Configurations with larger tiles, potentially better for B200
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'num_stages': 3, 'num_warps': 8}),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 32, 'num_stages': 3, 'num_warps': 8}),
triton.Config({'BLOCK_M': 256, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 3, 'num_warps': 8}),
# Configurations with deeper software pipelining
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 5, 'num_warps': 4}),
],
key=['M'],
)
@triton.jit
def gemm_kernel(
# Pointers to matrices
A, B, C,
# Matrix dimensions
M, N, K,
# Strides for matrices
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
# Meta-parameters
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr
):
"""
Triton kernel for GEMM: C = A @ B.T
A: [M, K]
B: [N, K]
C: [M, N]
"""
# -----------------------------------------------------------
# Map program ids (pids) to the block of C it should compute.
# This is a 1D launch grid, so we need to calculate the 2D block indices.
pid = tl.program_id(axis=0)
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
# ----------------------------------------------------------
# Create pointers for the first blocks of A and B.
# We will advance this pointer as we move in the K direction
# and accumulate pairs of tiles into C.
# Offsets for the M dimension of A and C
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
# Offsets for the N dimension of B and C
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Offsets for the K dimension
offs_k = tl.arange(0, BLOCK_K)
# Pointers to the first tile of A
# A is accessed as a [BLOCK_M, BLOCK_K] tile
a_ptrs = A + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
# Pointers to the first tile of B. We need to compute A @ B.T,
# so we load a [BLOCK_K, BLOCK_N] tile from B.T, which corresponds
# to B[n, k] elements.
b_ptrs = B + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)
# -----------------------------------------------------------
# Initialize accumulator.
# We accumulate in float32 for higher precision.
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# -----------------------------------------------------------
# Loop over the K dimension of A and B.
for k in range(0, tl.cdiv(K, BLOCK_K)):
# Load the next block of A and B.
# Masking is needed for the M dimension because M is variable.
# K=4096 and N=2048 are constants and multiples of the block sizes,
# so no masking is needed for them.
a = tl.load(a_ptrs, mask=offs_m[:, None] < M, other=0.0)
# --- FIX START ---
# The `other` argument requires a `mask`. Since no mask is needed for b,
# the `other` argument must be removed.
b = tl.load(b_ptrs)
# --- FIX END ---
# Perform the matrix multiplication of the tiles and accumulate the result.
accumulator += tl.dot(a, b)
# Advance the pointers to the next tile in the K dimension.
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# Cast the accumulator from float32 to float16 to match C's dtype.
c = accumulator.to(tl.float16)
# -----------------------------------------------------------
# Write the block of C back to global memory.
# Pointers to the destination block of C
offs_c = offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
c_ptrs = C + offs_c
# Masking is needed for the M dimension.
c_mask = (offs_m[:, None] < M)
tl.store(c_ptrs, c, mask=c_mask)
def gemm_n2048_k4096(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Wrapper function for the GEMM kernel C = A @ B.T.
Handles device management, tensor validation, and kernel launch.
Args:
A (torch.Tensor): A 2D tensor of shape [M, 4096] and dtype float16.
B (torch.Tensor): A 2D tensor of shape [2048, 4096] and dtype float16.
Returns:
torch.Tensor: The result tensor C of shape [M, 2048] and dtype float16.
"""
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This Triton kernel requires a GPU.")
# Store original devices to move results back later
device_A = A.device
# If inputs are on CPU, move to GPU. If on another device, raise error.
if A.device.type != 'cuda' or B.device.type != 'cuda':
if A.device.type == 'cpu' and B.device.type == 'cpu':
A = A.cuda()
B = B.cuda()
else:
raise RuntimeError("Input tensors must be on the same CUDA device or both on CPU.")
# --- Input Validation ---
assert A.is_contiguous(), "Input tensor A must be contiguous"
assert B.is_contiguous(), "Input tensor B must be contiguous"
assert A.dtype == torch.float16, "Input tensor A must have dtype float16"
assert B.dtype == torch.float16, "Input tensor B must have dtype float16"
assert A.dim() == 2 and B.dim() == 2, "Input tensors must be 2D"
M, K_A = A.shape
N, K_B = B.shape
# Check against the fixed dimensions from the specification
assert N == 2048, f"B.shape[0] must be 2048, but got {N}"
assert K_A == 4096, f"A.shape[1] must be 4096, but got {K_A}"
assert K_B == 4096, f"B.shape[1] must be 4096, but got {K_B}"
# --- Output Tensor ---
# Create the output tensor C on the same CUDA device.
C = torch.empty((M, N), device=A.device, dtype=torch.float16)
# --- Kernel Launch ---
# The grid is 1D, and we derive the 2D block indices inside the kernel.
# This is a common and effective pattern.
grid = lambda meta: (triton.cdiv(M, meta['BLOCK_M']) * triton.cdiv(N, meta['BLOCK_N']),)
gemm_kernel[grid](
A, B, C,
M, N, K_A,
A.stride(0), A.stride(1),
B.stride(0), B.stride(1),
C.stride(0), C.stride(1)
)
# --- Final Device Management ---
# Move the result back to the original device of the first input tensor.
if C.device != device_A:
C = C.to(device_A)
return C
def run(*args, **kwargs):
"""
Public entry point for the GEMM operation.
Handles both positional and keyword arguments for input tensors.
Args:
*args: Can contain two positional arguments (A, B).
**kwargs: Can contain two keyword arguments ('A', 'B').
Returns:
torch.Tensor: The result of the GEMM operation.
"""
if args and kwargs:
raise ValueError("Cannot provide both positional and keyword arguments for tensors.")
if 'A' in kwargs and 'B' in kwargs:
return gemm_n2048_k4096(kwargs['A'], kwargs['B'])
elif len(args) == 2:
return gemm_n2048_k4096(args[0], args[1])
else:
raise ValueError(
f"Invalid arguments. Expected 2 tensors, but got {len(args)} positional "
f"and {len(kwargs)} keyword args. Please provide (A, B) or A=..., B=..."
)scrolls · 196 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON