Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonvcx09o

gemini-2.5-pro_triton_vcx09o · gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 212 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-vcx09o?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16

Benchmark evidence

32 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n6144 k4096fp16 · [4, 4096]
NVIDIA B200
43.0µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [24, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [7, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [96, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [88, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [8, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [1, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [104, 4096]
NVIDIA B200
43.1µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [15, 4096]
NVIDIA B200
43.2µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [64, 4096]
NVIDIA B200
43.2µs
#3 of 6
2025-10-16
Show all 32 measurements ›
GEMM n6144 k4096fp16 · [80, 4096]
NVIDIA B200
43.2µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [128, 4096]
NVIDIA B200
43.4µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [248, 4096]
NVIDIA B200
43.6µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [152, 4096]
NVIDIA B200
43.6µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [200, 4096]
NVIDIA B200
43.6µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [216, 4096]
NVIDIA B200
43.7µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [224, 4096]
NVIDIA B200
43.7µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [136, 4096]
NVIDIA B200
43.7µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [192, 4096]
NVIDIA B200
43.8µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [160, 4096]
NVIDIA B200
43.8µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [232, 4096]
NVIDIA B200
43.8µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [168, 4096]
NVIDIA B200
43.9µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [240, 4096]
NVIDIA B200
43.9µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [208, 4096]
NVIDIA B200
44.0µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [256, 4096]
NVIDIA B200
44.0µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [176, 4096]
NVIDIA B200
44.0µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [184, 4096]
NVIDIA B200
44.0µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [40, 4096]
NVIDIA B200
47.2µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [972, 4096]
NVIDIA B200
76.6µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [2053, 4096]
NVIDIA B200
121.2µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [2379, 4096]
NVIDIA B200
143.6µs
#3 of 6
2025-10-16
GEMM n6144 k4096fp16 · [8192, 4096]
NVIDIA B200
424.1µs
#3 of 6
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:6c5bebb7a8d03bedb4cbce7cdab8a3d4c23681fc13d674f3da56932719fd9f18
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.

autotune@triton.autotune(
mmaaccumulator += tl.dot(a_tile, tl.trans(b_tile))

Kernel source

main.py212 lines
import torch
import triton
import triton.language as tl
import math

@triton.autotune(
    configs=[
        # Basic configurations
        triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
        triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'GROUP_SIZE_M': 8, 'num_stages': 4, 'num_warps': 4}),
        # Configurations with larger K block size
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 4}),
        triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 4}),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 8}),
        triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8, 'num_stages': 3, 'num_warps': 8}),
        # Potentially high-performing config for modern GPUs like B200
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8, 'num_stages': 2, 'num_warps': 8}),
    ],
    key=['M', 'N', 'K'],
)
@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,
    # Meta-parameters
    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.
    This kernel is optimized for large, constant N and K dimensions and a variable M dimension,
    targeting modern architectures like NVIDIA B200.

    - Tiling: The computation is broken down into tiles to maximize data reuse in fast memory.
    - Shared Memory: Tiles of A and B are loaded into shared memory to reduce global memory traffic.
    - Software Pipelining (`num_stages`): Overlaps memory access with computation to hide latency.
    - Grouped Scheduling (`GROUP_SIZE_M`): Encourages blocks that reuse data from matrix B to be
      scheduled on the same streaming multiprocessor, improving L2 cache hit rates.
    - FP32 Accumulator: Accumulation is done in `tl.float32` to maintain precision before
      storing the final `tl.float16` result.
    """
    # -----------------------------------------------------------
    # Map program ids to M and N dimensions using grouped scheduling
    pid = tl.program_id(axis=0)
    grid_m = tl.cdiv(M, BLOCK_SIZE_M)
    grid_n = tl.cdiv(N, BLOCK_SIZE_N)

    # Remap 1D program ID to 2D with grouping for better L2 cache locality
    width = GROUP_SIZE_M * grid_n
    group_id = pid // width
    group_size = tl.minimum(grid_m - group_id * GROUP_SIZE_M, GROUP_SIZE_M)
    
    pid_m = group_id * GROUP_SIZE_M + (pid % group_size)
    pid_n = (pid % width) // group_size

    # ----------------------------------------------------------
    # Create offsets for the C tile computed by this thread block
    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)
    
    # Create offsets for the K dimension
    offs_k = tl.arange(0, BLOCK_SIZE_K)

    # ----------------------------------------------------------
    # Initialize pointers to the input matrices A and B
    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 = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    # -----------------------------------------------------------
    # Loop over K in increments of BLOCK_SIZE_K
    for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
        # Load the next tile of A and B from global memory
        # Boundary checks are applied to handle cases where K is not a multiple of BLOCK_SIZE_K
        a_tile = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
        b_tile = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
        
        # Perform the matrix multiplication on the tiles.
        # We need to compute A @ B.T. `a_tile` is [BLOCK_SIZE_M, BLOCK_SIZE_K].
        # `b_tile` is loaded as [BLOCK_SIZE_N, BLOCK_SIZE_K], so we transpose it
        # to [BLOCK_SIZE_K, BLOCK_SIZE_N] for the dot product.
        accumulator += tl.dot(a_tile, tl.trans(b_tile))

        # Advance the 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_tile = accumulator.to(C.dtype.element_ty)

    # -----------------------------------------------------------
    # Write the result tile to global memory
    # Initialize pointers to the output matrix C
    c_ptrs = C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    
    # Create a mask to avoid out-of-bounds writes
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, c_tile, mask=c_mask)


def _validate_inputs(A, B):
    """Helper function to validate input tensor properties."""
    if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
        raise TypeError(f"Input must be torch.Tensor, got {type(A)}, {type(B)}")
    
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError(f"Input tensors must have dtype torch.float16, got {A.dtype}, {B.dtype}")
        
    # Check fixed dimensions N and K
    if A.shape[1] != 4096:
        raise ValueError(f"A.shape[1] must be 4096, but got {A.shape[1]}")
    if B.shape[0] != 6144:
        raise ValueError(f"B.shape[0] must be 6144, but got {B.shape[0]}")
    if B.shape[1] != 4096:
        raise ValueError(f"B.shape[1] must be 4096, but got {B.shape[1]}")
    
    if A.shape[1] != B.shape[1]:
        raise ValueError(f"Inner dimension K must match: A.shape[1]={A.shape[1]}, B.shape[1]={B.shape[1]}")


def _run_kernel(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    """
    Internal function to set up and launch the Triton kernel.
    Assumes inputs are already validated and on the correct GPU device.
    """
    A = A.contiguous()
    B = B.contiguous()
    
    M, K = A.shape
    N, _ = B.shape
    
    C = torch.empty((M, N), device=A.device, dtype=torch.float16)

    # Define the grid for the kernel launch using 1D grid for grouped scheduling
    grid = lambda META: (
        triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']),
    )
    
    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),
    )
    
    return C


def run(*args, **kwargs):
    """
    Entry point for the GEMM operation C = A @ B.T.

    This wrapper function handles device management, input validation,
    and kernel execution. It ensures that tensors are on the correct
    device (GPU) for the Triton kernel and that the result is moved back
    to the original device of the input tensors.

    Args:
        *args: Can be two positional arguments (A, B).
        **kwargs: Can be two keyword arguments (A=..., B=...).

    Returns:
        torch.Tensor: The result of the matrix multiplication, C.
    """
    if len(args) == 2 and not kwargs:
        A, B = args
    elif not args and 'A' in kwargs and 'B' in kwargs:
        A = kwargs.get('A')
        B = kwargs.get('B')
    else:
        raise ValueError("Invalid arguments. Use either positional (A, B) or keyword (A=tensor, B=tensor).")

    _validate_inputs(A, B)
    
    if not torch.cuda.is_available():
        raise RuntimeError("This kernel requires a CUDA-enabled GPU, but CUDA is not available.")
        
    original_device = A.device
    
    cuda_device = torch.device("cuda")
    if A.device.type != 'cuda' or B.device.type != 'cuda':
        try:
            A_gpu = A.to(cuda_device, non_blocking=True)
            B_gpu = B.to(cuda_device, non_blocking=True)
        except Exception as e:
            raise RuntimeError(f"Failed to move tensors to GPU: {e}")
    else:
        A_gpu = A
        B_gpu = B

    C_gpu = _run_kernel(A_gpu, B_gpu)
    
    if C_gpu.device != original_device:
        C_final = C_gpu.to(original_device, non_blocking=True)
    else:
        C_final = C_gpu
        
    return C_final
scrolls · 212 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON