Skip to content
KernelIndex
Search⌘K

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
GEMM n4096 k4096fp16 · [256, 4096]
NVIDIA B200
45.1µs
#5 of 9
2025-10-16
GEMM n4096 k4096fp16 · [240, 4096]
NVIDIA B200
45.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [248, 4096]
NVIDIA B200
45.1µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [232, 4096]
NVIDIA B200
45.1µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [224, 4096]
NVIDIA B200
45.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [192, 4096]
NVIDIA B200
45.1µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [208, 4096]
NVIDIA B200
45.2µs
#5 of 9
2025-10-16
GEMM n4096 k4096fp16 · [176, 4096]
NVIDIA B200
45.2µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [160, 4096]
NVIDIA B200
45.2µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [216, 4096]
NVIDIA B200
45.2µs
#4 of 7
2025-10-16
Show all 43 measurements ›
GEMM n4096 k4096fp16 · [128, 4096]
NVIDIA B200
45.2µs
#5 of 9
2025-10-16
GEMM n4096 k4096fp16 · [144, 4096]
NVIDIA B200
45.2µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [184, 4096]
NVIDIA B200
45.2µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [200, 4096]
NVIDIA B200
45.2µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [168, 4096]
NVIDIA B200
45.2µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [152, 4096]
NVIDIA B200
45.3µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [136, 4096]
NVIDIA B200
45.3µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [120, 4096]
NVIDIA B200
45.6µs
#5 of 8
2025-10-16
GEMM n4096 k4096fp16 · [112, 4096]
NVIDIA B200
46.5µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [972, 4096]
NVIDIA B200
46.8µs
#3 of 7
2025-10-16
GEMM n4096 k4096fp16 · [104, 4096]
NVIDIA B200
47.0µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [80, 4096]
NVIDIA B200
47.0µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [96, 4096]
NVIDIA B200
47.0µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [70, 4096]
NVIDIA B200
47.0µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [64, 4096]
NVIDIA B200
47.0µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [48, 4096]
NVIDIA B200
47.0µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [32, 4096]
NVIDIA B200
47.0µs
#3 of 8
2025-10-16
GEMM n4096 k4096fp16 · [40, 4096]
NVIDIA B200
47.0µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [24, 4096]
NVIDIA B200
47.1µs
#3 of 8
2025-10-16
GEMM n4096 k4096fp16 · [88, 4096]
NVIDIA B200
47.1µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [72, 4096]
NVIDIA B200
47.1µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [2, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [4, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [35, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [16, 4096]
NVIDIA B200
47.1µs
#3 of 8
2025-10-16
GEMM n4096 k4096fp16 · [56, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [7, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [8, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [15, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [1, 4096]
NVIDIA B200
47.1µs
#4 of 8
2025-10-16
GEMM n4096 k4096fp16 · [2379, 4096]
NVIDIA B200
109.0µs
#3 of 7
2025-10-16
GEMM n4096 k4096fp16 · [2053, 4096]
NVIDIA B200
117.8µs
#4 of 7
2025-10-16
GEMM n4096 k4096fp16 · [8192, 4096]
NVIDIA B200
258.7µs
#3 of 8
2025-10-16

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.

mmaaccumulator += tl.dot(a, b_t)
num-warps = 8num_warps = 8
tile-k = 128BLOCK_M, BLOCK_N, BLOCK_K = 128, 256, 64

Kernel 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