gemini-2.5-pro / tritonr3ccri
gemini-2.5-pro_triton_r3ccri · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 195 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-r3ccri?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
43 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 43 measurements ›Showing all 43 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:4cf689d3a5d70342971d232511c50eede8654b1c336995584c83120b9e990538
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.
mma
accumulator += tl.dot(a, b_t)num-warps = 8
num_warps = 8tile-k = 128
BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 64Kernel source
main.py195 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gemm_kernel(
A, B, C,
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,
GROUP_M: tl.constexpr
):
"""
Triton kernel for GEMM: C = A @ B.T, optimized for B200.
This kernel computes matrix multiplication for A[M, K] and B[N, K] to produce C[M, N].
It is optimized for a B200-class GPU by using:
- Large tile sizes (BLOCK_M, BLOCK_N) to increase arithmetic intensity.
- A large BLOCK_K to improve data reuse from shared memory.
- Software pipelining managed by the Triton compiler to hide memory latency.
- FP32 accumulation for numerical stability before converting to FP16 output.
The operation is C = A @ B.T, which translates to C[m, n] = sum_k(A[m, k] * B[n, k]).
This means we load contiguous blocks from both A and B.
"""
# -----------------------------------------------------------
# Map program ids to M, N blocks
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
# For grouped launch, calculate the specific block indices
pid_group = pid // num_pid_n
pid_n = pid % num_pid_n
# Each group of blocks works on a contiguous region of M
group_start_m = pid_group * GROUP_M
pid_m = group_start_m + (tl.program_id(axis=1) % GROUP_M)
# ----------------------------------------------------------
# Create pointers for the first blocks of A and B.
# We use block pointers to efficiently load tiles from global memory.
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)
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
# The accumulator is in float32 to prevent precision loss
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# -----------------------------------------------------------
# Main loop over the K dimension
for k in range(0, tl.cdiv(K, BLOCK_K)):
# Boundary checks for K
k_remaining = K - k * BLOCK_K
k_mask = offs_k[None, :] < k_remaining
# Load the next block of A and B from global memory
# Masking is applied to handle cases where K is not a multiple of BLOCK_K,
# and where M is not a multiple of BLOCK_M.
a_mask = (offs_m[:, None] < M) & k_mask
b_mask = k_mask # N is constant and a multiple of BLOCK_N, so no N mask needed for B load
a = tl.load(a_ptrs, mask=a_mask, other=0.0)
b = tl.load(b_ptrs, mask=b_mask, other=0.0)
# The operation is A @ B.T. We load a tile from B [BLOCK_N, BLOCK_K]
# and transpose it to [BLOCK_K, BLOCK_N] before the dot product.
# tl.trans is efficient for register-level transposition.
b_t = tl.trans(b)
# Perform the matrix multiplication
accumulator += tl.dot(a, b_t)
# Advance the pointers to the next K block
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# -----------------------------------------------------------
# Cast accumulator to float16 and write back to C
c = accumulator.to(tl.float16)
# Create pointers to the C matrix
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
c_ptrs = C + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
# Create a mask to avoid out-of-bounds writes
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
def gemm_n4096_k4096(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Wrapper function for the GEMM kernel: C = A @ B.T.
Args:
A (torch.Tensor): A tensor of shape [M, 4096] and dtype float16.
B (torch.Tensor): A tensor of shape [4096, 4096] and dtype float16.
Returns:
torch.Tensor: The result tensor C of shape [M, 4096] and dtype float16.
"""
# --- Dimention and Dtype Checks ---
assert A.shape[1] == 4096, f"A.shape[1] must be 4096, but is {A.shape[1]}"
assert B.shape[0] == 4096, f"B.shape[0] must be 4096, but is {B.shape[0]}"
assert B.shape[1] == 4096, f"B.shape[1] must be 4096, but is {B.shape[1]}"
assert A.dtype == torch.float16, f"A.dtype must be float16, but is {A.dtype}"
assert B.dtype == torch.float16, f"B.dtype must be float16, but is {B.dtype}"
M, K = A.shape
N, K_check = B.shape
# --- Output Tensor ---
# The output tensor is created on the same device as the inputs.
C = torch.empty((M, N), device=A.device, dtype=A.dtype)
# --- Kernel Configuration ---
# Configuration chosen for B200-like architecture.
# BLOCK_M, BLOCK_N: Large tile sizes to maximize compute-to-memory ratio.
# BLOCK_K: Balances shared memory usage and data reuse.
# num_warps: Uses 8 warps (256 threads) per block for high occupancy.
# GROUP_M: Groups thread blocks to improve L2 cache locality for the M-dimension.
BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 64
GROUP_M = 8
num_warps = 8
# --- Grid Calculation ---
# The grid is 2D, but we launch it as a 1D grid of "groups" and a 1D grid of blocks within a group.
grid_m = triton.cdiv(M, BLOCK_M)
grid_n = triton.cdiv(N, BLOCK_N)
# We group blocks along the M dimension to improve L2 cache hit rate
grid = (triton.cdiv(grid_m, GROUP_M) * grid_n, GROUP_M)
# --- Kernel Launch ---
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),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
GROUP_M=GROUP_M,
num_warps=num_warps,
# Triton's compiler automatically handles software pipelining.
# For this kernel structure, a num_stages of 3 or 4 is typical.
)
return C
def run(*args, **kwargs):
"""
Public entry point for the GEMM operation.
This function handles device management and calls the Triton kernel.
It accepts tensors 'A' and 'B' via args or kwargs.
"""
if 'A' in kwargs and 'B' in kwargs:
A = kwargs['A']
B = kwargs['B']
elif len(args) == 2:
A, B = args
else:
raise ValueError("Please provide tensors 'A' and 'B' as arguments or keyword arguments.")
if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
raise TypeError("Inputs 'A' and 'B' must be torch.Tensors.")
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("Triton requires a CUDA-enabled GPU, but CUDA is not available.")
# Store original device to return the result on the same device
original_device = A.device
# Determine the target GPU device. If any input is on a GPU, use it.
# Otherwise, move inputs to the default CUDA device.
gpu_device = next((t.device for t in [A, B] if t.is_cuda), torch.device('cuda'))
A_gpu = A.to(gpu_device)
B_gpu = B.to(gpu_device)
# --- Execute and Return ---
C_gpu = gemm_n4096_k4096(A_gpu, B_gpu)
return C_gpu.to(original_device)
scrolls · 195 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON