gemini-2.5-pro / tritonmryn73
gemini-2.5-pro_triton_mryn73 · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 189 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-mryn73?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:0d5dcb2293f898c6628143f5b871a5c3f995bc70f9c743a9a983e49beec34e3f
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)num-warps = 8
num_warps = 8stages = 3
num_stages = 3 # Software pipelining stagetile-k = 64
BLOCK_K = 64 # K=4096 is divisible by 64tile-m = 128
BLOCK_M = 128tile-n = 256
BLOCK_N = 256 # N=28672 is divisible by 256Kernel source
main.py189 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gemm_kernel(
# Pointers to matrices
A, B, C,
# Matrix dimensions
M, N, K,
# Strides
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,
GROUP_M: tl.constexpr
):
"""
Triton kernel for GEMM C = A @ B.T, optimized for B200-class GPUs.
This kernel computes the matrix multiplication of A ([M, K]) and the transpose of B ([N, K]),
storing the result in C ([M, N]).
Key optimizations:
- Tiling: The computation is broken down into smaller blocks (BLOCK_M, BLOCK_N, BLOCK_K)
that fit into shared memory, maximizing data reuse.
- Software Pipelining: `num_stages` in the host code enables pre-fetching of data,
hiding global memory latency.
- Grouped Scheduling (GROUP_M): Improves L2 cache locality for large M dimensions by
processing chunks of A against all of B before moving to the next chunk.
- Vectorized Loads/Stores: Triton's compiler automatically vectorizes memory operations.
- Optimized for N=28672, K=4096: The block sizes are chosen such that no bounds checking
is needed for the N and K dimensions, simplifying the inner loop.
"""
# -----------------------------------------------------------
# Grid and program ID calculation with grouped scheduling
pid = tl.program_id(axis=0)
# Total number of program instances along M and N axes
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
# Number of programs in a group
num_pid_in_group = GROUP_M * num_pid_n
# ID of the group this program belongs to
group_id = pid // num_pid_in_group
# Row-major order within a group for better L2 cache locality
first_pid_m = group_id * GROUP_M
pid_in_group = pid % num_pid_in_group
# ID of the M-tile and N-tile within the group
pid_m = first_pid_m + (pid_in_group // num_pid_n)
pid_n = pid_in_group % num_pid_n
# Guard against out-of-bounds work items when M is not a multiple of BLOCK_M*GROUP_M
if pid_m >= num_pid_m:
return
# ----------------------------------------------------------
# Pointers to the first element of the blocks
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)
# Pointers for the A block [BLOCK_M, BLOCK_K]
a_ptrs = A + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
# Pointers for the B block, loaded as [BLOCK_K, BLOCK_N] to match dot product
# This corresponds to accessing B[n, k] for the matmul A @ B.T
b_ptrs = B + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)
# -----------------------------------------------------------
# Main loop over K-dimension
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, tl.cdiv(K, BLOCK_K)):
# Load A and B blocks from global memory
# Boundary check for M is needed as M is variable.
# No checks needed for N and K as they are constants divisible by block sizes.
m_mask = offs_m[:, None] < M
a = tl.load(a_ptrs, mask=m_mask, other=0.0)
b = tl.load(b_ptrs) # No mask needed for B
# Matrix multiplication using Tensor Cores
accumulator += tl.dot(a, b)
# Advance pointers to the next K block
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# Cast accumulator to the output dtype
c = accumulator.to(tl.float16)
# -----------------------------------------------------------
# Write back the result to C
# Pointers to the C block
c_ptrs = C + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
# Store the result, masking for the variable M dimension
store_mask = offs_m[:, None] < M
tl.store(c_ptrs, c, mask=store_mask)
def run(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Wrapper function for the GEMM operation C = A @ B.T.
Handles device management, kernel launching, and returns the result on the
original device of the input tensors.
Args:
A (torch.Tensor): A 2D tensor of shape [M, 4096] and dtype float16.
B (torch.Tensor): A 2D tensor of shape [28672, 4096] and dtype float16.
Returns:
torch.Tensor: The result C of the matrix multiplication, with shape [M, 28672]
and dtype float16, on the same device as the input tensors.
"""
# ---- Validation ----
# Validate dimensions and dtypes based on the problem specification
K_DIM = 4096
N_DIM = 28672
if A.shape[1] != K_DIM:
raise ValueError(f"Input A must have K={K_DIM}, but got shape {A.shape}")
if B.shape[0] != N_DIM or B.shape[1] != K_DIM:
raise ValueError(f"Input B must have shape [{N_DIM}, {K_DIM}], but got shape {B.shape}")
if A.dtype != torch.float16:
raise TypeError(f"Input A must be float16, but got {A.dtype}")
if B.dtype != torch.float16:
raise TypeError(f"Input B must be float16, but got {B.dtype}")
# ---- Device Management ----
original_device = A.device
if original_device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
device = torch.device("cuda")
A_gpu = A.to(device)
B_gpu = B.to(device)
elif original_device.type == 'cuda':
device = original_device
A_gpu = A
B_gpu = B
else:
raise TypeError(f"Unsupported device type: {original_device.type}. Only 'cpu' and 'cuda' are supported.")
# ---- Kernel Execution ----
M, K = A_gpu.shape
N, _ = B_gpu.shape
# Allocate output tensor on the GPU
C = torch.empty((M, N), device=device, dtype=torch.float16)
# Kernel configuration optimized for B200-like architectures
# These parameters use large tile sizes to maximize compute utilization and hide memory latency.
BLOCK_M = 128
BLOCK_N = 256 # N=28672 is divisible by 256
BLOCK_K = 64 # K=4096 is divisible by 64
GROUP_M = 8 # Grouping for L2 cache locality
num_warps = 8
num_stages = 3 # Software pipelining stage
# The grid is 1D, and the kernel partitions it into a 2D grid with grouping
grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N), )
# Launch the kernel
gemm_kernel[grid](
A_gpu, B_gpu, C,
M, N, K,
A_gpu.stride(0), A_gpu.stride(1),
B_gpu.stride(0), B_gpu.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,
num_stages=num_stages
)
# ---- Return Result ----
# Move the result back to the original device of the inputs
return C.to(original_device)
scrolls · 189 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON