Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton5iu7uf

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

25 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n5120 k2048fp16 · [34, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [64, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [32, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [63, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [1, 2048]
NVIDIA B200
24.6µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [5, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [25, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [93, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [172, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
Show all 25 measurements ›
GEMM n5120 k2048fp16 · [6, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [4, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [128, 2048]
NVIDIA B200
24.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [289, 2048]
NVIDIA B200
24.7µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [17, 2048]
NVIDIA B200
24.7µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [2, 2048]
NVIDIA B200
24.7µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [16, 2048]
NVIDIA B200
24.7µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [492, 2048]
NVIDIA B200
26.3µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [952, 2048]
NVIDIA B200
40.5µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8828, 2048]
NVIDIA B200
190.5µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [11006, 2048]
NVIDIA B200
226.6µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12853, 2048]
NVIDIA B200
267.3µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12251, 2048]
NVIDIA B200
273.6µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [14915, 2048]
NVIDIA B200
309.3µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [16294, 2048]
NVIDIA B200
345.5µ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:5df2577f34e44188492725e911f0e23756bb00a8169f059c9f232208b13c1bd4
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, tl.trans(b))

Kernel source

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

@triton.autotune(
    configs=[
        # Basic configurations
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'num_warps': 4, 'num_stages': 2}, num_ctas=1),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32, 'num_warps': 8, 'num_stages': 2}, num_ctas=1),
        triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32, 'num_warps': 8, 'num_stages': 2}, num_ctas=1),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 64, 'num_warps': 4, 'num_stages': 3}, num_ctas=1),
        triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'num_warps': 4, 'num_stages': 3}, num_ctas=1),
        # Configurations potentially good for B200 with larger compute/memory resources
        triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'num_warps': 8, 'num_stages': 3}, num_ctas=1),
        triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'num_warps': 8, 'num_stages': 3}, num_ctas=1),
    ],
    key=['M'],
)
@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,
):
    """
    Computes C = A @ B.T
    A is of shape (M, K)
    B is of shape (N, K)
    C is of shape (M, N)
    This kernel is optimized for a matrix multiplication where the second matrix (B)
    is transposed. Both A and B are expected to be row-major.
    The kernel is structured to perform coalesced loads from both A and B.
    The transpose operation is handled by `tl.trans` on the register-loaded tile of B
    before the `tl.dot` operation. This approach relies on the compiler to efficiently
    schedule the transpose and dot instructions and is effective on modern GPUs with
    large caches like B200.
    """
    # -----------------------------------------------------------
    # Map program ids to M and N dimensions.
    pid = tl.program_id(axis=0)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    # ----------------------------------------------------------
    # Create pointers for the first blocks of A and B.
    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.
    # The accumulator is in float32 to maintain precision.
    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    # -----------------------------------------------------------
    # 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 the case where M is not a multiple of BLOCK_SIZE_M.
        # Since N and K are constants and our block sizes divide them, masks for N and K
        # are not strictly necessary but are kept for generality. The compiler will optimize
        # them out if possible.
        a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K), other=0.0)
        b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_k[None, :] < K), other=0.0)
        
        # The operation is C = A @ B.T, which translates to C[m,n] = sum_k A[m,k] * B[n,k].
        # Our loaded tile `a` is (BLOCK_M, BLOCK_K) and `b` is (BLOCK_N, BLOCK_K).
        # We need to compute dot(a, b.T). tl.trans(b) makes it (BLOCK_K, BLOCK_N).
        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 output dtype and write back to C.
    C_out = accumulator.to(tl.float16)

    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_out, mask=c_mask)

def run(*args, **kwargs):
    """
    Wrapper function for the GEMM kernel, providing a user-friendly interface
    and handling all device management.

    Args:
        A (torch.Tensor): The first input tensor of shape [M, K].
        B (torch.Tensor): The second input tensor of shape [N, K].
        Can be passed as positional or keyword arguments.

    Returns:
        torch.Tensor: The output tensor C of shape [M, N], on the same device as the input A.
    """
    # --- Argument parsing ---
    if len(args) == 2:
        A, B = args
    elif 'A' in kwargs and 'B' in kwargs:
        A = kwargs['A']
        B = kwargs['B']
    else:
        raise ValueError("Inputs 'A' and 'B' must be provided either as positional or keyword arguments.")

    # --- Shape and DType validation ---
    assert A.dtype == torch.float16, f"Input A must be float16, but got {A.dtype}"
    assert B.dtype == torch.float16, f"Input B must be float16, but got {B.dtype}"
    
    M, K_A = A.shape
    N, K_B = B.shape
    
    # Constants from the spec
    spec_N, spec_K = 5120, 2048
    
    assert K_A == spec_K, f"A.shape[1] must be {spec_K}, but got {K_A}"
    assert K_B == spec_K, f"B.shape[1] must be {spec_K}, but got {K_B}"
    assert N == spec_N, f"B.shape[0] must be {spec_N}, but got {N}"

    # --- Device Management ---
    if not torch.cuda.is_available():
        raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
    
    original_device = A.device
    device = torch.device("cuda")

    # Move tensors to GPU if they are not already there
    A_gpu = A.to(device, non_blocking=True) if A.device != device else A
    B_gpu = B.to(device, non_blocking=True) if B.device != device else B

    # Ensure inputs are contiguous for optimal memory access
    A_gpu = A_gpu.contiguous()
    B_gpu = B_gpu.contiguous()

    # --- Output Tensor Allocation ---
    C = torch.empty((M, N), device=device, dtype=torch.float16)

    # --- Grid Definition ---
    # We use a 1D grid to simplify launching and autotuning, especially for the dynamic M dimension.
    # The kernel then internally maps the 1D program ID to 2D (M, N) block coordinates.
    grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']),)

    # --- Kernel Launch ---
    gemm_kernel[grid](
        A_gpu, B_gpu, C,
        M, N, spec_K,
        A_gpu.stride(0), A_gpu.stride(1),
        B_gpu.stride(0), B_gpu.stride(1),
        C.stride(0), C.stride(1),
    )

    # --- Result Handling ---
    # Move the result back to the original device of the input tensor
    return C.to(original_device, non_blocking=True)
scrolls · 165 lines total

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

Best evidence level for this revision: reported

JSON