Skip to content
KernelIndex
Search⌘K

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
GEMM n256 k7168fp16 · [901, 7168]
NVIDIA B200
67.7µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [80, 7168]
NVIDIA B200
69.6µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [56, 7168]
NVIDIA B200
70.2µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [63, 7168]
NVIDIA B200
70.2µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [55, 7168]
NVIDIA B200
71.3µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [58, 7168]
NVIDIA B200
71.5µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [54, 7168]
NVIDIA B200
71.6µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [57, 7168]
NVIDIA B200
71.6µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [32, 7168]
NVIDIA B200
71.6µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [53, 7168]
NVIDIA B200
71.6µs
#3 of 7
2025-10-16
Show all 17 measurements ›
GEMM n256 k7168fp16 · [15, 7168]
NVIDIA B200
71.7µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [14, 7168]
NVIDIA B200
71.7µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [16, 7168]
NVIDIA B200
72.2µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [4, 7168]
NVIDIA B200
73.0µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [1, 7168]
NVIDIA B200
73.1µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [11948, 7168]
NVIDIA B200
87.1µs
#3 of 7
2025-10-16
GEMM n256 k7168fp16 · [14104, 7168]
NVIDIA B200
90.0µs
#3 of 7
2025-10-16

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.

mmaaccumulator = 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