Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonnekk4o

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

29 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n2048 k4096fp16 · [492, 4096]
NVIDIA B200
39.0µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [952, 4096]
NVIDIA B200
40.8µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [969, 4096]
NVIDIA B200
40.9µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [172, 4096]
NVIDIA B200
41.0µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [289, 4096]
NVIDIA B200
41.0µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [93, 4096]
NVIDIA B200
41.0µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [63, 4096]
NVIDIA B200
41.0µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [128, 4096]
NVIDIA B200
41.1µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [64, 4096]
NVIDIA B200
41.1µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [34, 4096]
NVIDIA B200
41.3µs
#2 of 7
2025-10-16
Show all 29 measurements ›
GEMM n2048 k4096fp16 · [32, 4096]
NVIDIA B200
41.9µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [5, 4096]
NVIDIA B200
42.1µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [25, 4096]
NVIDIA B200
42.2µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [2, 4096]
NVIDIA B200
42.4µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15, 4096]
NVIDIA B200
42.4µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16, 4096]
NVIDIA B200
42.5µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [17, 4096]
NVIDIA B200
42.6µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [4, 4096]
NVIDIA B200
42.6µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [1, 4096]
NVIDIA B200
42.6µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8, 4096]
NVIDIA B200
42.6µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [6, 4096]
NVIDIA B200
42.7µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8828, 4096]
NVIDIA B200
156.7µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11006, 4096]
NVIDIA B200
189.6µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11938, 4096]
NVIDIA B200
212.7µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12251, 4096]
NVIDIA B200
216.9µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12853, 4096]
NVIDIA B200
228.7µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [14915, 4096]
NVIDIA B200
265.8µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15813, 4096]
NVIDIA B200
271.2µs
#2 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16294, 4096]
NVIDIA B200
271.8µs
#2 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:cfdc62bc0b2ab6a656f937835ddf9debeecfeb72049ec5db049a4ae9f3799d0c
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, b)

Kernel source

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

@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 2, 'num_warps': 4}),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 32, 'num_stages': 2, 'num_warps': 4}),
        # Configurations with larger tiles, potentially better for B200
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'num_stages': 3, 'num_warps': 8}),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 32, 'num_stages': 3, 'num_warps': 8}),
        triton.Config({'BLOCK_M': 256, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 3, 'num_warps': 8}),
        # Configurations with deeper software pipelining
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64, 'num_stages': 4, 'num_warps': 4}),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 32, 'num_stages': 5, 'num_warps': 4}),
    ],
    key=['M'],
)
@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_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr
):
    """
    Triton kernel for GEMM: C = A @ B.T
    A: [M, K]
    B: [N, K]
    C: [M, N]
    """
    # -----------------------------------------------------------
    # Map program ids (pids) to the block of C it should compute.
    # This is a 1D launch grid, so we need to calculate the 2D block indices.
    pid = tl.program_id(axis=0)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    # ----------------------------------------------------------
    # Create pointers for the first blocks of A and B.
    # We will advance this pointer as we move in the K direction
    # and accumulate pairs of tiles into C.
    
    # Offsets for the M dimension of A and C
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    # Offsets for the N dimension of B and C
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    # Offsets for the K dimension
    offs_k = tl.arange(0, BLOCK_K)

    # Pointers to the first tile of A
    # A is accessed as a [BLOCK_M, BLOCK_K] tile
    a_ptrs = A + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
    
    # Pointers to the first tile of B. We need to compute A @ B.T,
    # so we load a [BLOCK_K, BLOCK_N] tile from B.T, which corresponds
    # to B[n, k] elements.
    b_ptrs = B + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)
    
    # -----------------------------------------------------------
    # Initialize accumulator.
    # We accumulate in float32 for higher precision.
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # -----------------------------------------------------------
    # Loop over the K dimension of A and B.
    for k in range(0, tl.cdiv(K, BLOCK_K)):
        # Load the next block of A and B.
        # Masking is needed for the M dimension because M is variable.
        # K=4096 and N=2048 are constants and multiples of the block sizes,
        # so no masking is needed for them.
        a = tl.load(a_ptrs, mask=offs_m[:, None] < M, other=0.0)
        # --- FIX START ---
        # The `other` argument requires a `mask`. Since no mask is needed for b,
        # the `other` argument must be removed.
        b = tl.load(b_ptrs)
        # --- FIX END ---
        
        # Perform the matrix multiplication of the tiles and accumulate the result.
        accumulator += tl.dot(a, b)
        
        # Advance the pointers to the next tile in the K dimension.
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    # Cast the accumulator from float32 to float16 to match C's dtype.
    c = accumulator.to(tl.float16)

    # -----------------------------------------------------------
    # Write the block of C back to global memory.
    # Pointers to the destination block of C
    offs_c = offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    c_ptrs = C + offs_c
    
    # Masking is needed for the M dimension.
    c_mask = (offs_m[:, None] < M)
    tl.store(c_ptrs, c, mask=c_mask)


def gemm_n2048_k4096(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    """
    Wrapper function for the GEMM kernel C = A @ B.T.
    Handles device management, tensor validation, and kernel launch.

    Args:
        A (torch.Tensor): A 2D tensor of shape [M, 4096] and dtype float16.
        B (torch.Tensor): A 2D tensor of shape [2048, 4096] and dtype float16.

    Returns:
        torch.Tensor: The result tensor C of shape [M, 2048] and dtype float16.
    """
    # --- Device Management ---
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available. This Triton kernel requires a GPU.")

    # Store original devices to move results back later
    device_A = A.device
    # If inputs are on CPU, move to GPU. If on another device, raise error.
    if A.device.type != 'cuda' or B.device.type != 'cuda':
        if A.device.type == 'cpu' and B.device.type == 'cpu':
            A = A.cuda()
            B = B.cuda()
        else:
            raise RuntimeError("Input tensors must be on the same CUDA device or both on CPU.")
    
    # --- Input Validation ---
    assert A.is_contiguous(), "Input tensor A must be contiguous"
    assert B.is_contiguous(), "Input tensor B must be contiguous"
    assert A.dtype == torch.float16, "Input tensor A must have dtype float16"
    assert B.dtype == torch.float16, "Input tensor B must have dtype float16"
    assert A.dim() == 2 and B.dim() == 2, "Input tensors must be 2D"

    M, K_A = A.shape
    N, K_B = B.shape

    # Check against the fixed dimensions from the specification
    assert N == 2048, f"B.shape[0] must be 2048, but got {N}"
    assert K_A == 4096, f"A.shape[1] must be 4096, but got {K_A}"
    assert K_B == 4096, f"B.shape[1] must be 4096, but got {K_B}"
    
    # --- Output Tensor ---
    # Create the output tensor C on the same CUDA device.
    C = torch.empty((M, N), device=A.device, dtype=torch.float16)

    # --- Kernel Launch ---
    # The grid is 1D, and we derive the 2D block indices inside the kernel.
    # This is a common and effective pattern.
    grid = lambda meta: (triton.cdiv(M, meta['BLOCK_M']) * triton.cdiv(N, meta['BLOCK_N']),)
    
    gemm_kernel[grid](
        A, B, C,
        M, N, K_A,
        A.stride(0), A.stride(1),
        B.stride(0), B.stride(1),
        C.stride(0), C.stride(1)
    )

    # --- Final Device Management ---
    # Move the result back to the original device of the first input tensor.
    if C.device != device_A:
        C = C.to(device_A)

    return C

def run(*args, **kwargs):
    """
    Public entry point for the GEMM operation.
    Handles both positional and keyword arguments for input tensors.

    Args:
        *args: Can contain two positional arguments (A, B).
        **kwargs: Can contain two keyword arguments ('A', 'B').
    
    Returns:
        torch.Tensor: The result of the GEMM operation.
    """
    if args and kwargs:
        raise ValueError("Cannot provide both positional and keyword arguments for tensors.")
    
    if 'A' in kwargs and 'B' in kwargs:
        return gemm_n2048_k4096(kwargs['A'], kwargs['B'])
    elif len(args) == 2:
        return gemm_n2048_k4096(args[0], args[1])
    else:
        raise ValueError(
            f"Invalid arguments. Expected 2 tensors, but got {len(args)} positional "
            f"and {len(kwargs)} keyword args. Please provide (A, B) or A=..., B=..."
        )
scrolls · 196 lines total

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

Best evidence level for this revision: reported

JSON