Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonmryn73

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-mryn73?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 n28672 k4096fp16 · [120, 4096]
NVIDIA B200
68.1µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [104, 4096]
NVIDIA B200
68.4µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [88, 4096]
NVIDIA B200
68.5µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [72, 4096]
NVIDIA B200
68.7µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [24, 4096]
NVIDIA B200
68.7µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [40, 4096]
NVIDIA B200
68.8µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [56, 4096]
NVIDIA B200
68.8µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8, 4096]
NVIDIA B200
68.8µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [7, 4096]
NVIDIA B200
68.9µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2, 4096]
NVIDIA B200
68.9µs
#5 of 8
2025-10-16
Show all 43 measurements ›
GEMM n28672 k4096fp16 · [128, 4096]
NVIDIA B200
68.9µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [112, 4096]
NVIDIA B200
68.9µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [80, 4096]
NVIDIA B200
69.0µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [96, 4096]
NVIDIA B200
69.1µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [70, 4096]
NVIDIA B200
69.1µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [64, 4096]
NVIDIA B200
69.2µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [16, 4096]
NVIDIA B200
69.3µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [48, 4096]
NVIDIA B200
69.3µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [15, 4096]
NVIDIA B200
69.3µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [35, 4096]
NVIDIA B200
69.3µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [1, 4096]
NVIDIA B200
69.4µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [4, 4096]
NVIDIA B200
69.4µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [32, 4096]
NVIDIA B200
69.4µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [232, 4096]
NVIDIA B200
109.6µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [256, 4096]
NVIDIA B200
109.7µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [216, 4096]
NVIDIA B200
109.7µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [152, 4096]
NVIDIA B200
109.8µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [136, 4096]
NVIDIA B200
109.8µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [200, 4096]
NVIDIA B200
109.9µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [168, 4096]
NVIDIA B200
109.9µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [184, 4096]
NVIDIA B200
110.0µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [240, 4096]
NVIDIA B200
110.0µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [224, 4096]
NVIDIA B200
110.2µs
#4 of 8
2025-10-16
GEMM n28672 k4096fp16 · [208, 4096]
NVIDIA B200
110.4µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [192, 4096]
NVIDIA B200
110.6µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [144, 4096]
NVIDIA B200
110.6µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [160, 4096]
NVIDIA B200
110.6µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [176, 4096]
NVIDIA B200
110.7µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [248, 4096]
NVIDIA B200
113.5µs
#5 of 8
2025-10-16
GEMM n28672 k4096fp16 · [972, 4096]
NVIDIA B200
357.5µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2053, 4096]
NVIDIA B200
719.8µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [2379, 4096]
NVIDIA B200
814.6µs
#6 of 8
2025-10-16
GEMM n28672 k4096fp16 · [8192, 4096]
NVIDIA B200
2.57ms
#6 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:0d5dcb2293f898c6628143f5b871a5c3f995bc70f9c743a9a983e49beec34e3f
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)
num-warps = 8num_warps = 8
stages = 3num_stages = 3 # Software pipelining stage
tile-k = 64BLOCK_K = 64 # K=4096 is divisible by 64
tile-m = 128BLOCK_M = 128
tile-n = 256BLOCK_N = 256 # N=28672 is divisible by 256

Kernel source

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

@triton.jit
def gemm_kernel(
    # Pointers to matrices
    A, B, C,
    # Matrix dimensions
    M, N, K,
    # Strides
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    # Meta-parameters
    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-class GPUs.
    
    This kernel computes the matrix multiplication of A ([M, K]) and the transpose of B ([N, K]),
    storing the result in C ([M, N]).
    
    Key optimizations:
    - Tiling: The computation is broken down into smaller blocks (BLOCK_M, BLOCK_N, BLOCK_K)
      that fit into shared memory, maximizing data reuse.
    - Software Pipelining: `num_stages` in the host code enables pre-fetching of data,
      hiding global memory latency.
    - Grouped Scheduling (GROUP_M): Improves L2 cache locality for large M dimensions by
      processing chunks of A against all of B before moving to the next chunk.
    - Vectorized Loads/Stores: Triton's compiler automatically vectorizes memory operations.
    - Optimized for N=28672, K=4096: The block sizes are chosen such that no bounds checking
      is needed for the N and K dimensions, simplifying the inner loop.
    """
    # -----------------------------------------------------------
    # Grid and program ID calculation with grouped scheduling
    pid = tl.program_id(axis=0)
    
    # Total number of program instances along M and N axes
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    
    # Number of programs in a group
    num_pid_in_group = GROUP_M * num_pid_n
    # ID of the group this program belongs to
    group_id = pid // num_pid_in_group
    
    # Row-major order within a group for better L2 cache locality
    first_pid_m = group_id * GROUP_M
    pid_in_group = pid % num_pid_in_group
    
    # ID of the M-tile and N-tile within the group
    pid_m = first_pid_m + (pid_in_group // num_pid_n)
    pid_n = pid_in_group % num_pid_n

    # Guard against out-of-bounds work items when M is not a multiple of BLOCK_M*GROUP_M
    if pid_m >= num_pid_m:
        return

    # ----------------------------------------------------------
    # Pointers to the first element of the blocks
    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)

    # Pointers for the A block [BLOCK_M, BLOCK_K]
    a_ptrs = A + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
    
    # Pointers for the B block, loaded as [BLOCK_K, BLOCK_N] to match dot product
    # This corresponds to accessing B[n, k] for the matmul A @ B.T
    b_ptrs = B + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)
    
    # -----------------------------------------------------------
    # Main loop over K-dimension
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for k in range(0, tl.cdiv(K, BLOCK_K)):
        # Load A and B blocks from global memory
        # Boundary check for M is needed as M is variable.
        # No checks needed for N and K as they are constants divisible by block sizes.
        m_mask = offs_m[:, None] < M
        
        a = tl.load(a_ptrs, mask=m_mask, other=0.0)
        b = tl.load(b_ptrs) # No mask needed for B
        
        # Matrix multiplication using Tensor Cores
        accumulator += tl.dot(a, b)

        # Advance pointers to the next K block
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    # Cast accumulator to the output dtype
    c = accumulator.to(tl.float16)

    # -----------------------------------------------------------
    # Write back the result to C
    # Pointers to the C block
    c_ptrs = C + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
    
    # Store the result, masking for the variable M dimension
    store_mask = offs_m[:, None] < M
    tl.store(c_ptrs, c, mask=store_mask)


def run(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    """
    Wrapper function for the GEMM operation C = A @ B.T.

    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, 4096] and dtype float16.
        B (torch.Tensor): A 2D tensor of shape [28672, 4096] and dtype float16.

    Returns:
        torch.Tensor: The result C of the matrix multiplication, with shape [M, 28672]
                      and dtype float16, on the same device as the input tensors.
    """
    # ---- Validation ----
    # Validate dimensions and dtypes based on the problem specification
    K_DIM = 4096
    N_DIM = 28672
    if A.shape[1] != K_DIM:
        raise ValueError(f"Input A must have K={K_DIM}, but got shape {A.shape}")
    if B.shape[0] != N_DIM or B.shape[1] != K_DIM:
        raise ValueError(f"Input B must have shape [{N_DIM}, {K_DIM}], but got shape {B.shape}")
    if A.dtype != torch.float16:
        raise TypeError(f"Input A must be float16, but got {A.dtype}")
    if B.dtype != torch.float16:
        raise TypeError(f"Input B must be float16, but got {B.dtype}")

    # ---- Device Management ----
    original_device = A.device
    
    if original_device.type == 'cpu':
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
        device = torch.device("cuda")
        A_gpu = A.to(device)
        B_gpu = B.to(device)
    elif original_device.type == 'cuda':
        device = original_device
        A_gpu = A
        B_gpu = B
    else:
        raise TypeError(f"Unsupported device type: {original_device.type}. Only 'cpu' and 'cuda' are supported.")

    # ---- Kernel Execution ----
    M, K = A_gpu.shape
    N, _ = B_gpu.shape

    # Allocate output tensor on the GPU
    C = torch.empty((M, N), device=device, dtype=torch.float16)

    # Kernel configuration optimized for B200-like architectures
    # These parameters use large tile sizes to maximize compute utilization and hide memory latency.
    BLOCK_M = 128
    BLOCK_N = 256  # N=28672 is divisible by 256
    BLOCK_K = 64   # K=4096 is divisible by 64
    GROUP_M = 8    # Grouping for L2 cache locality
    num_warps = 8
    num_stages = 3 # Software pipelining stage

    # The grid is 1D, and the kernel partitions it into a 2D grid with grouping
    grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N), )

    # Launch the kernel
    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),
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        GROUP_M=GROUP_M,
        num_warps=num_warps,
        num_stages=num_stages
    )

    # ---- Return Result ----
    # Move the result back to the original device of the inputs
    return C.to(original_device)
scrolls · 189 lines total

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

Best evidence level for this revision: reported

JSON