Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritondc28mj

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

14 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Fused add RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
6.64µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
7.28µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
7.31µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
7.42µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
7.42µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
7.42µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
7.47µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
7.70µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
8.16µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
43.1µs
#1 of 8
2025-10-16
Show all 14 measurements ›
Fused add RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
51.0µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
54.0µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
63.5µs
#1 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
63.5µs
#1 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:14fb30ecff672694407ade17e0e63783a7fe2acab06591fca11b38f132c9af44
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(
num-warps = 4triton.Config({'BLOCK_SIZE_H': 1024}, num_warps=4),

Kernel source

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

@triton.autotune(
    configs=[
        triton.Config({'BLOCK_SIZE_H': 1024}, num_warps=4),
        triton.Config({'BLOCK_SIZE_H': 2048}, num_warps=8),
        triton.Config({'BLOCK_SIZE_H': 4096}, num_warps=16),
    ],
    key=['HIDDEN_SIZE'],
)
@triton.jit
def _fused_add_rmsnorm_h4096_kernel(
    # Pointers to tensors
    hidden_states_ptr,
    residual_ptr,
    weight_ptr,
    output_ptr,
    # Other parameters
    batch_size,
    # Constants
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE_H: tl.constexpr,
):
    """
    Triton kernel for fused add and RMSNorm.
    Each program instance processes a single row of the input tensors.
    """
    # Get the row index for the current program
    row_idx = tl.program_id(0)

    # --- Pointers for the current row ---
    # Note: This assumes inputs are contiguous, which is a common case.
    row_hidden_states_ptr = hidden_states_ptr + row_idx * HIDDEN_SIZE
    row_residual_ptr = residual_ptr + row_idx * HIDDEN_SIZE
    row_output_ptr = output_ptr + row_idx * HIDDEN_SIZE
    # Weight pointer is the same for all rows
    weight_base_ptr = weight_ptr

    # --- Pass 1: Calculate the sum of squares for RMSNorm ---
    # Accumulator for the variance, initialized to zero.
    # We use float32 for precision.
    var_accumulator = tl.zeros([1], dtype=tl.float32)

    # Iterate over the hidden dimension in blocks of size BLOCK_SIZE_H
    for h_offset in range(0, HIDDEN_SIZE, BLOCK_SIZE_H):
        # Create a vector of offsets for the current block
        h_offsets = h_offset + tl.arange(0, BLOCK_SIZE_H)
        # Create a mask to handle the last block if HIDDEN_SIZE is not a multiple of BLOCK_SIZE_H
        mask = h_offsets < HIDDEN_SIZE

        # Load one block of hidden_states and residual
        hidden_states_block = tl.load(row_hidden_states_ptr + h_offsets, mask=mask, other=0.0)
        residual_block = tl.load(row_residual_ptr + h_offsets, mask=mask, other=0.0)

        # Perform the addition: x = hidden_states + residual
        # Cast to float32 for high-precision computation
        x = hidden_states_block.to(tl.float32) + residual_block.to(tl.float32)

        # Accumulate the sum of squares of x
        var_accumulator += tl.sum(x * x, axis=0)

    # After iterating through all blocks, compute the mean and inverse RMS
    mean_var = var_accumulator / HIDDEN_SIZE
    inv_rms = tl.rsqrt(mean_var + EPS)

    # --- Pass 2: Apply normalization and scaling, and store the result ---
    # Re-iterate over the hidden dimension to apply the calculated inv_rms
    for h_offset in range(0, HIDDEN_SIZE, BLOCK_SIZE_H):
        h_offsets = h_offset + tl.arange(0, BLOCK_SIZE_H)
        mask = h_offsets < HIDDEN_SIZE

        # Reload the input blocks for this pass
        hidden_states_block = tl.load(row_hidden_states_ptr + h_offsets, mask=mask, other=0.0)
        residual_block = tl.load(row_residual_ptr + h_offsets, mask=mask, other=0.0)

        # Recompute x, same as in Pass 1
        x = hidden_states_block.to(tl.float32) + residual_block.to(tl.float32)

        # Load the corresponding block of weights
        weight_block = tl.load(weight_base_ptr + h_offsets, mask=mask, other=0.0)

        # Apply the normalization and scaling
        normalized_x = x * inv_rms
        output_block = normalized_x * weight_block.to(tl.float32)

        # Convert the final result back to bfloat16 and store it in the output tensor
        tl.store(row_output_ptr + h_offsets, output_block.to(tl.bfloat16), mask=mask)


def fused_add_rmsnorm_h4096(hidden_states: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
    """
    Wrapper function for the fused_add_rmsnorm_h4096 Triton kernel.

    Args:
        hidden_states (torch.Tensor): Input tensor of shape [batch_size, 4096] and dtype bfloat16.
        residual (torch.Tensor): Residual tensor of shape [batch_size, 4096] and dtype bfloat16.
        weight (torch.Tensor): Weight tensor of shape [4096] and dtype bfloat16.

    Returns:
        torch.Tensor: The output tensor of shape [batch_size, 4096] and dtype bfloat16.
    """
    # --- Device Management ---
    if not torch.cuda.is_available():
        raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")

    original_device = hidden_states.device
    is_cpu = original_device.type == 'cpu'

    # Move tensors to GPU if they are on CPU
    if is_cpu:
        hidden_states = hidden_states.cuda()
        residual = residual.cuda()
        weight = weight.cuda()

    # --- Input Validation ---
    batch_size, hidden_size = hidden_states.shape
    
    if hidden_size != 4096:
        raise ValueError(f"hidden_size must be 4096, but got {hidden_size}")
    if hidden_states.shape != residual.shape:
        raise ValueError("hidden_states and residual must have the same shape.")
    if weight.shape != (hidden_size,):
        raise ValueError(f"weight must have shape [{hidden_size}], but got {weight.shape}")
    
    expected_dtype = torch.bfloat16
    if hidden_states.dtype != expected_dtype or residual.dtype != expected_dtype or weight.dtype != expected_dtype:
        raise TypeError(f"All input tensors must have dtype {expected_dtype}.")

    # --- Kernel Launch ---
    # Allocate output tensor on the same device as the inputs
    output = torch.empty_like(hidden_states)

    # The grid is 1D, with one program instance per row in the batch
    grid = (batch_size,)

    _fused_add_rmsnorm_h4096_kernel[grid](
        hidden_states,
        residual,
        weight,
        output,
        batch_size,
        HIDDEN_SIZE=hidden_size,
        EPS=1e-5,
    )

    # --- Finalization ---
    # Move output back to the original device if necessary
    if is_cpu:
        output = output.cpu()

    return output


def run(*args, **kwargs):
    """
    Public entry point for the kernel. Handles flexible argument passing.
    This function allows calling with either positional or keyword arguments.
    """
    # A robust way to handle both *args and **kwargs by forwarding them
    # to the main function.
    if args:
        return fused_add_rmsnorm_h4096(*args, **kwargs)
    else:
        return fused_add_rmsnorm_h4096(**kwargs)
scrolls · 169 lines total

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

Best evidence level for this revision: reported

JSON