gemini-2.5-pro / tritonq84sir
gemini-2.5-pro_triton_q84sir · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 191 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-q84sir?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:5e901cd6435e296fd6aebba5ccf70f4ec94edf49e9000df93881cd09a2b6443f
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))num-warps = 8
num_warps=8,stages = 3
num_stages=3tile-k = 64
BLOCK_SIZE_K=64,tile-m = 128
BLOCK_SIZE_M=128,tile-n = 128
BLOCK_SIZE_N=128,Kernel source
main.py191 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_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 optimized for B200.
A is (M, K), B is (N, K), C is (M, N).
This kernel computes C[m, n] = sum_k(A[m, k] * B[n, k]).
Tuning and Strategy:
- Tiling: The problem is tiled into blocks of size (BLOCK_SIZE_M, BLOCK_SIZE_N) for the output C.
The K dimension is processed in blocks of BLOCK_SIZE_K.
- Grouping: Thread blocks are grouped (GROUP_SIZE_M) along the M dimension to improve L2 cache reuse
for the B matrix. Blocks within a group are executed by a single SM or cluster, keeping the
relevant parts of B in L2 cache. This is achieved by remapping a 1D grid launch.
- Pipelining: Software pipelining (num_stages in the launcher) is used to hide the latency of loading
data from global memory. While one block of A and B is being processed, the next is being fetched.
- Data Types: FP32 is used for accumulation to maintain precision and avoid overflow, before
converting the final result back to FP16.
- B200 Optimizations: This kernel uses large block sizes (128x128) and a high thread count per block
(8 warps = 256 threads) to fully utilize the compute resources of modern architectures like B200.
The `tl.dot` instruction maps directly to Tensor Core operations.
"""
# -----------------------------------------------------------
# Map program ids to tiles
# This kernel uses a 1D launch grid and remaps it to a 2D grid with grouping.
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
# Remap to a 2D grid with grouping for L2 cache locality
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
# Within the group, blocks are arranged column-major
pid_m = first_pid_m + (pid % group_size)
pid_n = (pid % num_pid_in_group) // group_size
# ----------------------------------------------------------
# Create pointers for the first blocks of A and B.
# rm/rn will be block-level offsets, offs_m/n will be thread-level offsets
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)
# Pointers to the start of the first tile for this program instance
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)
# -----------------------------------------------------------
# Accumulator initialization
# Use float32 for higher precision and to avoid overflow
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# -----------------------------------------------------------
# Main loop over the K dimension
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
# Load the next block of A and B
# Masking is applied to handle cases where K is not a multiple of BLOCK_SIZE_K
# and for the variable M dimension.
a_mask = (offs_m[:, None] < M) & (offs_k[None, :] + k * BLOCK_SIZE_K < K)
b_mask = (offs_n[:, None] < N) & (offs_k[None, :] + k * BLOCK_SIZE_K < K)
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, which means we need to compute dot(A, B.T)
# We load a tile from B of shape [BLOCK_N, BLOCK_K]. Triton's tl.dot requires
# the second operand to be of shape [K, N]. The load pattern for `b` is already
# effectively transposing it for memory access efficiency.
# So we use tl.trans to get B into the right shape for tl.dot.
accumulator += tl.dot(a, tl.trans(b))
# Advance 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 = accumulator.to(tl.float16)
# -----------------------------------------------------------
# Write back the block of the output matrix C
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
c_ptrs = C + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
def run(*args, **kwargs):
"""
Public entry point for the gemm_n4096_k14336 operation.
This function serves as a complete wrapper, handling device management,
tensor validation, kernel launching, and result handling.
Args:
A (torch.Tensor): The first input tensor of shape [M, 14336] and dtype float16.
B (torch.Tensor): The second input tensor of shape [4096, 14336] and dtype float16.
Returns:
torch.Tensor: The output tensor C of shape [M, 4096] and dtype float16,
on the same device as the input tensors.
"""
# -----------------------------------------------------------
# Argument parsing
# -----------------------------------------------------------
if 'A' in kwargs and 'B' in kwargs:
A = kwargs['A']
B = kwargs['B']
elif len(args) == 2:
A, B = args
else:
raise ValueError("Expected 2 positional arguments (A, B) or keyword arguments 'A' and 'B'.")
# -----------------------------------------------------------
# Constants and validation
# -----------------------------------------------------------
N, K = 4096, 14336
assert A.shape[1] == K, f"Input A has invalid K dimension: {A.shape[1]}, expected {K}"
assert B.shape[0] == N, f"Input B has invalid N dimension: {B.shape[0]}, expected {N}"
assert B.shape[1] == K, f"Input B has invalid K dimension: {B.shape[1]}, expected {K}"
assert A.dtype == torch.float16, f"Input A must be float16, but is {A.dtype}"
assert B.dtype == torch.float16, f"Input B must be float16, but is {B.dtype}"
M = A.shape[0]
# -----------------------------------------------------------
# Device management
# -----------------------------------------------------------
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
original_device = A.device
is_cpu = original_device.type == 'cpu'
if is_cpu:
A_gpu = A.cuda()
B_gpu = B.cuda()
else:
A_gpu = A
B_gpu = B
# -----------------------------------------------------------
# Kernel launch
# -----------------------------------------------------------
# Allocate output tensor
C = torch.empty((M, N), device=A_gpu.device, dtype=torch.float16)
# Grid computation
grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), )
# Kernel call
# Using a single, well-tuned configuration for B200.
# In a real-world scenario, this would be autotuned.
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),
# --- Kernel meta-parameters ---
BLOCK_SIZE_M=128,
BLOCK_SIZE_N=128,
BLOCK_SIZE_K=64,
GROUP_SIZE_M=8,
# num_stages and num_warps are passed to the Triton compiler
# For B200, 8 warps and 3+ stages are good starting points
num_warps=8,
num_stages=3
)
# -----------------------------------------------------------
# Final device management
# -----------------------------------------------------------
if is_cpu:
return C.to(original_device)
else:
return C
scrolls · 191 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON