gemini-2.5-pro_triton_015737
gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 225 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-015737?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
17 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 17 measurements ›Showing all 17 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:43efecfea85d665416fabdb14a376fd453bc73dec569ec34dcce09389a969a48
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, tl.trans(b), accumulator, allow_tf32=True)Kernel source
main.py225 lines
import torch
import triton
import triton.language as tl
import math
# Triton Kernel for GEMM: C = A @ B.T
@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,
# Tile sizes
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
# Grouping for L2 cache performance
GROUP_SIZE_M: tl.constexpr
):
"""
Computes C = A @ B.T where A is [M, K] and B is [N, K].
This kernel is optimized for modern NVIDIA GPUs like B200.
- Tiling strategy is chosen for the given fixed N and K dimensions.
- Grouped block ordering is used to improve L2 cache hit rate for the B matrix.
- Software pipelining is enabled via num_stages to hide memory latency.
"""
# -----------------------------------------------------------
# Map program ids to M and N blocks
# -----------------------------------------------------------
pid = tl.program_id(axis=0)
# Grid dimensions
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
# Grouping programs for better L2 cache locality
# Programs are grouped together along the M dimension to reuse B matrix tiles
num_pids_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pids_in_group
# M and N block indices for this program
first_pid_m = group_id * GROUP_SIZE_M
group_size = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % group_size)
pid_n = (pid % num_pids_in_group) // group_size
# ----------------------------------------------------------
# Create pointers for the first blocks of A and B.
# We will advance these pointers as we loop over K.
# ----------------------------------------------------------
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)
offs_k = tl.arange(0, BLOCK_SIZE_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.
# Accumulator holds the C tile, computed in float32 for precision.
# -----------------------------------------------------------
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# -----------------------------------------------------------
# Main loop over the K dimension
# -----------------------------------------------------------
# Loop until the K dimension is fully processed.
# tl.cdiv is used to handle the case where K is not a multiple of BLOCK_SIZE_K,
# though for this specific problem K (7168) is a multiple of BLOCK_SIZE_K (64).
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
# Load the next block of A and B from global memory.
# Masking is applied to handle the variable M dimension.
a_mask = offs_m[:, None] < M
a = tl.load(a_ptrs, mask=a_mask, other=0.0)
# For N and K, masking is not needed because they are fixed and perfectly
# divisible by their respective block sizes.
b = tl.load(b_ptrs)
# Perform the matrix multiplication on the loaded tiles.
# The result is accumulated in float32.
# The B matrix tile is transposed implicitly by tl.dot.
accumulator = tl.dot(a, tl.trans(b), accumulator, allow_tf32=True)
# Advance the pointers to the next K block.
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += BLOCK_SIZE_K * stride_bk
# -----------------------------------------------------------
# Write the result to the output matrix C
# -----------------------------------------------------------
# Cast the accumulator from float32 to the output dtype (float16).
c = accumulator.to(C.dtype.element_ty)
# Create pointers to the C matrix and apply masks for storing.
c_ptrs = C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
c_mask_m = offs_m[:, None] < M
c_mask_n = offs_n[None, :] < N # This mask is always true but is good practice
c_mask = c_mask_m & c_mask_n
tl.store(c_ptrs, c, mask=c_mask)
def gemm_n256_k7168(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
Computes the matrix multiplication C = A @ B.T using a Triton kernel.
This function is a wrapper that 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, 7168] and dtype float16.
B (torch.Tensor): A 2D tensor of shape [256, 7168] and dtype float16.
Returns:
torch.Tensor: The result C of shape [M, 256] and dtype float16.
"""
# --- Device Management ---
# Ensure CUDA is available
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
# Preserve the original device of the input tensor to return the output on the same device
original_device = A.device
# Move tensors to GPU. If they are already on the correct GPU, this is a no-op.
device = torch.device('cuda')
A = A.to(device)
B = B.to(device)
# --- Input Validation ---
# Check tensor dimensions and dtypes
assert A.dim() == 2 and B.dim() == 2, "Input tensors must be 2D"
assert A.dtype == torch.float16, "Input tensor A must be of dtype float16"
assert B.dtype == torch.float16, "Input tensor B must be of dtype float16"
# Get matrix dimensions
M, K = A.shape
N, K_check = B.shape
# Validate against the kernel's fixed dimensions
assert N == 256, f"Dimension N of B must be 256, but got {N}"
assert K == 7168, f"Dimension K of A must be 7168, but got {K}"
assert K == K_check, f"Inner dimension K of A and B must match, but got {K} and {K_check}"
# --- Output Tensor Allocation ---
# Allocate the output tensor on the GPU
C = torch.empty((M, N), device=device, dtype=torch.float16)
# --- Kernel Launch Configuration ---
# This configuration is optimized for B200-class GPUs.
# BLOCK_SIZE_M: Tile size for the M dimension.
# BLOCK_SIZE_N: Tile size for the N dimension. Set to N to simplify loops.
# BLOCK_SIZE_K: Tile size for the K dimension.
# GROUP_SIZE_M: Number of M-blocks to group together for L2 cache reuse.
# num_warps: Number of warps per thread block.
# num_stages: Number of pipeline stages for loading from global memory.
# A strong configuration for Hopper/Blackwell architecture
config = {
'BLOCK_SIZE_M': 128,
'BLOCK_SIZE_N': 256,
'BLOCK_SIZE_K': 64,
'GROUP_SIZE_M': 8,
'num_warps': 8,
'num_stages': 3
}
# Define the launch grid
# The grid is 1D, where each program computes one C tile.
def grid(meta):
# The multiplication by cdiv(N, BLOCK_SIZE_N) is technically `* 1` here,
# but it's the general form for a 2D-tiled problem.
return (triton.cdiv(M, meta['BLOCK_SIZE_M']) * triton.cdiv(N, meta['BLOCK_SIZE_N']), )
# --- Kernel Execution ---
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_SIZE_M=config['BLOCK_SIZE_M'],
BLOCK_SIZE_N=config['BLOCK_SIZE_N'],
BLOCK_SIZE_K=config['BLOCK_SIZE_K'],
GROUP_SIZE_M=config['GROUP_SIZE_M'],
num_warps=config['num_warps'],
num_stages=config['num_stages']
)
# --- Return Result ---
# Move the result tensor back to the original device of the inputs
return C.to(original_device)
def run(*args, **kwargs):
"""
Public entry point for the GEMM operation.
This function handles flexible argument parsing (args and kwargs) and
delegates to the main implementation.
Args can be provided as `run(A, B)` or kwargs as `run(A=A_tensor, B=B_tensor)`.
"""
A = kwargs.get('A')
B = kwargs.get('B')
if A is None:
if len(args) > 0:
A = args[0]
else:
raise ValueError("Missing required input tensor 'A'")
if B is None:
if len(args) > 1:
B = args[1]
else:
raise ValueError("Missing required input tensor 'B'")
return gemm_n256_k7168(A, B)
scrolls · 225 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON