gemini-2.5-pro / tritonvcx09o
gemini-2.5-pro_triton_vcx09o · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 212 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-vcx09o?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
32 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 32 measurements ›Showing all 32 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:6c5bebb7a8d03bedb4cbce7cdab8a3d4c23681fc13d674f3da56932719fd9f18
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.py212 lines
import torch
import triton
import triton.language as tl
import math
@triton.autotune(
configs=[
# Basic configurations
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
# Configurations with larger K block size
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 4}),
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 8}),
triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 8}),
# Potentially high-performing config for modern GPUs like B200
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8, 'num_stages': 2, 'num_warps': 8}),
],
key=['M', 'N', 'K'],
)
@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_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
):
"""
Triton kernel for GEMM C = A @ B.T.
This kernel is optimized for large, constant N and K dimensions and a variable M dimension,
targeting modern architectures like NVIDIA B200.
- Tiling: The computation is broken down into tiles to maximize data reuse in fast memory.
- Shared Memory: Tiles of A and B are loaded into shared memory to reduce global memory traffic.
- Software Pipelining (`num_stages`): Overlaps memory access with computation to hide latency.
- Grouped Scheduling (`GROUP_SIZE_M`): Encourages blocks that reuse data from matrix B to be
scheduled on the same streaming multiprocessor, improving L2 cache hit rates.
- FP32 Accumulator: Accumulation is done in `tl.float32` to maintain precision before
storing the final `tl.float16` result.
"""
# -----------------------------------------------------------
# Map program ids to M and N dimensions using grouped scheduling
pid = tl.program_id(axis=0)
grid_m = tl.cdiv(M, BLOCK_SIZE_M)
grid_n = tl.cdiv(N, BLOCK_SIZE_N)
# Remap 1D program ID to 2D with grouping for better L2 cache locality
width = GROUP_SIZE_M * grid_n
group_id = pid // width
group_size = tl.minimum(grid_m - group_id * GROUP_SIZE_M, GROUP_SIZE_M)
pid_m = group_id * GROUP_SIZE_M + (pid % group_size)
pid_n = (pid % width) // group_size
# ----------------------------------------------------------
# Create offsets for the C tile computed by this thread block
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)
# Create offsets for the K dimension
offs_k = tl.arange(0, BLOCK_SIZE_K)
# ----------------------------------------------------------
# Initialize pointers to the input matrices A and B
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 with zeros
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# -----------------------------------------------------------
# Loop over K in increments of BLOCK_SIZE_K
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
# Load the next tile of A and B from global memory
# Boundary checks are applied to handle cases where K is not a multiple of BLOCK_SIZE_K
a_tile = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
b_tile = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
# Perform the matrix multiplication on the tiles.
# We need to compute A @ B.T. `a_tile` is [BLOCK_SIZE_M, BLOCK_SIZE_K].
# `b_tile` is loaded as [BLOCK_SIZE_N, BLOCK_SIZE_K], so we transpose it
# to [BLOCK_SIZE_K, BLOCK_SIZE_N] for the dot product.
accumulator += tl.dot(a_tile, tl.trans(b_tile))
# Advance the 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
c_tile = accumulator.to(C.dtype.element_ty)
# -----------------------------------------------------------
# Write the result tile to global memory
# Initialize pointers to the output matrix C
c_ptrs = C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
# Create a mask to avoid out-of-bounds writes
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, c_tile, mask=c_mask)
def _validate_inputs(A, B):
"""Helper function to validate input tensor properties."""
if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
raise TypeError(f"Input must be torch.Tensor, got {type(A)}, {type(B)}")
if A.dtype != torch.float16 or B.dtype != torch.float16:
raise TypeError(f"Input tensors must have dtype torch.float16, got {A.dtype}, {B.dtype}")
# Check fixed dimensions N and K
if A.shape[1] != 4096:
raise ValueError(f"A.shape[1] must be 4096, but got {A.shape[1]}")
if B.shape[0] != 6144:
raise ValueError(f"B.shape[0] must be 6144, but got {B.shape[0]}")
if B.shape[1] != 4096:
raise ValueError(f"B.shape[1] must be 4096, but got {B.shape[1]}")
if A.shape[1] != B.shape[1]:
raise ValueError(f"Inner dimension K must match: A.shape[1]={A.shape[1]}, B.shape[1]={B.shape[1]}")
def _run_kernel(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Internal function to set up and launch the Triton kernel.
Assumes inputs are already validated and on the correct GPU device.
"""
A = A.contiguous()
B = B.contiguous()
M, K = A.shape
N, _ = B.shape
C = torch.empty((M, N), device=A.device, dtype=torch.float16)
# Define the grid for the kernel launch using 1D grid for grouped scheduling
grid = lambda META: (
triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']),
)
gemm_kernel[grid](
A, B, C,
M, N, K,
A.stride(0), A.stride(1),
B.stride(0), B.stride(1),
C.stride(0), C.stride(1),
)
return C
def run(*args, **kwargs):
"""
Entry point for the GEMM operation C = A @ B.T.
This wrapper function handles device management, input validation,
and kernel execution. It ensures that tensors are on the correct
device (GPU) for the Triton kernel and that the result is moved back
to the original device of the input tensors.
Args:
*args: Can be two positional arguments (A, B).
**kwargs: Can be two keyword arguments (A=..., B=...).
Returns:
torch.Tensor: The result of the matrix multiplication, C.
"""
if len(args) == 2 and not kwargs:
A, B = args
elif not args and 'A' in kwargs and 'B' in kwargs:
A = kwargs.get('A')
B = kwargs.get('B')
else:
raise ValueError("Invalid arguments. Use either positional (A, B) or keyword (A=tensor, B=tensor).")
_validate_inputs(A, B)
if not torch.cuda.is_available():
raise RuntimeError("This kernel requires a CUDA-enabled GPU, but CUDA is not available.")
original_device = A.device
cuda_device = torch.device("cuda")
if A.device.type != 'cuda' or B.device.type != 'cuda':
try:
A_gpu = A.to(cuda_device, non_blocking=True)
B_gpu = B.to(cuda_device, non_blocking=True)
except Exception as e:
raise RuntimeError(f"Failed to move tensors to GPU: {e}")
else:
A_gpu = A
B_gpu = B
C_gpu = _run_kernel(A_gpu, B_gpu)
if C_gpu.device != original_device:
C_final = C_gpu.to(original_device, non_blocking=True)
else:
C_final = C_gpu
return C_final
scrolls · 212 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON