Skip to content
KernelIndex
Search⌘K

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
GEMM n4096 k14336fp16 · [112, 14336]
NVIDIA B200
128.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [120, 14336]
NVIDIA B200
128.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [128, 14336]
NVIDIA B200
129.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [104, 14336]
NVIDIA B200
129.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [96, 14336]
NVIDIA B200
129.3µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [168, 14336]
NVIDIA B200
129.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [160, 14336]
NVIDIA B200
129.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [136, 14336]
NVIDIA B200
129.5µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [232, 14336]
NVIDIA B200
129.6µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [176, 14336]
NVIDIA B200
129.7µs
#2 of 6
2025-10-16
Show all 43 measurements ›
GEMM n4096 k14336fp16 · [88, 14336]
NVIDIA B200
130.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [152, 14336]
NVIDIA B200
130.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [144, 14336]
NVIDIA B200
130.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [80, 14336]
NVIDIA B200
130.5µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [72, 14336]
NVIDIA B200
131.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [70, 14336]
NVIDIA B200
131.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [64, 14336]
NVIDIA B200
131.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [192, 14336]
NVIDIA B200
131.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [200, 14336]
NVIDIA B200
131.6µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [56, 14336]
NVIDIA B200
131.9µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [208, 14336]
NVIDIA B200
132.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [240, 14336]
NVIDIA B200
132.1µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [184, 14336]
NVIDIA B200
132.3µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [48, 14336]
NVIDIA B200
133.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [40, 14336]
NVIDIA B200
133.6µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [216, 14336]
NVIDIA B200
133.7µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [35, 14336]
NVIDIA B200
134.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [224, 14336]
NVIDIA B200
134.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [32, 14336]
NVIDIA B200
134.6µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [256, 14336]
NVIDIA B200
135.1µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [24, 14336]
NVIDIA B200
135.5µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [248, 14336]
NVIDIA B200
135.5µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [15, 14336]
NVIDIA B200
136.1µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [8, 14336]
NVIDIA B200
136.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [16, 14336]
NVIDIA B200
136.4µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [2, 14336]
NVIDIA B200
136.6µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [1, 14336]
NVIDIA B200
136.9µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [7, 14336]
NVIDIA B200
137.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [4, 14336]
NVIDIA B200
137.0µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [972, 14336]
NVIDIA B200
172.2µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [2053, 14336]
NVIDIA B200
330.5µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [2379, 14336]
NVIDIA B200
398.7µs
#2 of 6
2025-10-16
GEMM n4096 k14336fp16 · [8192, 14336]
NVIDIA B200
1.03ms
#2 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: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.

mmaaccumulator += tl.dot(a, tl.trans(b))
num-warps = 8num_warps=8,
stages = 3num_stages=3
tile-k = 64BLOCK_SIZE_K=64,
tile-m = 128BLOCK_SIZE_M=128,
tile-n = 128BLOCK_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