Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton0owo26

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-0owo26?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
RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
9.30µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
10.6µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
12.0µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
12.5µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
12.9µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
13.1µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
13.2µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
13.2µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
13.7µs
#5 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
28.8µs
#1 of 6
2025-10-16
Show all 14 measurements ›
RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
33.9µs
#1 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
35.7µs
#1 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
41.4µs
#1 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
41.7µs
#1 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:c5d6dcefab0abc44d57188a1f74800b43548a83d568cb3033b4e32a30480650d
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.

num-warps = 8num_warps=8, # A good value for B200 to ensure high occupancy

Kernel source

main.py175 lines
import torch
import triton
import triton.language as tl
import math
import inspect

#
# Triton Kernel
#
@triton.jit
def rmsnorm_h4096_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    stride_hidden_batch,
    stride_output_batch,
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE_H: tl.constexpr,
):
    """
    Triton kernel for RMS Normalization, optimized for a fixed hidden_size of 4096.

    This kernel is designed for B200 and performs the operation in a single fused pass
    from the perspective of HBM. It reads hidden_states and weights once, and writes
    the output once. The intermediate calculations (sum of squares, inv_rms) are
    done in registers/SRAM.

    Grid:
      - The grid is 1D, with one program instance per batch element (row).
    
    Memory Access:
      - Each program loads an entire row of hidden_states (4096 elements).
      - It also loads the entire weight vector (4096 elements).
      - Accesses are contiguous and coalesced.

    Computation:
      - All intermediate computations are performed in float32 for precision.
      - The final result is cast back to bfloat16 before storing.
    """
    # Get the index for the current batch element.
    batch_idx = tl.program_id(axis=0)

    # --- Pointer Setup ---
    # Pointers to the start of the current row for input and output tensors.
    row_hidden_states_ptr = hidden_states_ptr + batch_idx * stride_hidden_batch
    row_output_ptr = output_ptr + batch_idx * stride_output_batch

    # Offsets for accessing the hidden dimension.
    h_offsets = tl.arange(0, BLOCK_SIZE_H)
    
    # --- Pass 1: Compute Inverse RMS ---
    # Load the entire row of hidden_states for the current batch element.
    # Since HIDDEN_SIZE is fixed at 4096, we load it in one go.
    x_bf16 = tl.load(row_hidden_states_ptr + h_offsets)
    x_fp32 = x_bf16.to(tl.float32)

    # Calculate the sum of squares. tl.sum performs a parallel reduction
    # within the thread block.
    sum_sq = tl.sum(x_fp32 * x_fp32, axis=0)
    
    # Calculate the mean and the inverse root mean square.
    mean_sq = sum_sq / HIDDEN_SIZE
    inv_rms = tl.rsqrt(mean_sq + EPS)

    # --- Pass 2: Apply Normalization and Store ---
    # Load the weight vector.
    w_bf16 = tl.load(weight_ptr + h_offsets)
    w_fp32 = w_bf16.to(tl.float32)
    
    # Apply the RMS normalization and scale by the weights.
    # The original x_fp32 values should still be in registers.
    output_fp32 = x_fp32 * inv_rms * w_fp32
    
    # Cast the final result back to bfloat16 and store it.
    output_bf16 = output_fp32.to(tl.bfloat16)
    tl.store(row_output_ptr + h_offsets, output_bf16)


#
# Wrapper Function
#
def _rmsnorm_h4096_impl(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
    """
    Wrapper for the RMSNorm Triton kernel.

    This function handles device management, shape/dtype validation,
    grid computation, and kernel invocation.
    """
    # --- Device Management & Validation ---
    if not torch.cuda.is_available():
        raise RuntimeError("This implementation requires a CUDA-enabled GPU.")
    
    original_device = hidden_states.device
    is_cpu_input = original_device.type == 'cpu'

    # Move tensors to GPU if they are on CPU
    if is_cpu_input:
        hidden_states = hidden_states.cuda()
        weight = weight.cuda()
    
    # Validate shapes and dtypes
    batch_size, hidden_size = hidden_states.shape
    
    if hidden_size != 4096:
        raise ValueError(f"This kernel is specialized for hidden_size=4096, but got {hidden_size}.")
    if weight.shape != (hidden_size,):
        raise ValueError(f"Weight shape mismatch. Expected ({hidden_size},) but got {weight.shape}.")
    if hidden_states.dtype != torch.bfloat16:
        raise TypeError(f"Input 'hidden_states' must be of dtype torch.bfloat16, but got {hidden_states.dtype}.")
    if weight.dtype != torch.bfloat16:
        raise TypeError(f"Input 'weight' must be of dtype torch.bfloat16, but got {weight.dtype}.")

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

    # Define the grid for the kernel launch. One program per batch element.
    grid = (batch_size,)
    
    # Constants for the kernel.
    HIDDEN_SIZE = 4096
    EPS = 1e-5

    # Launch the Triton kernel.
    rmsnorm_h4096_kernel[grid](
        hidden_states,
        weight,
        output,
        hidden_states.stride(0),
        output.stride(0),
        HIDDEN_SIZE=HIDDEN_SIZE,
        EPS=EPS,
        BLOCK_SIZE_H=HIDDEN_SIZE, # Process the entire row in one block
        num_warps=8, # A good value for B200 to ensure high occupancy
    )

    # --- Final Device Management ---
    # If the original input was on CPU, move the result back to CPU.
    if is_cpu_input:
        output = output.to(original_device)

    return output

#
# Public Entry Point
#
def run(*args, **kwargs):
    """
    Public entry point for the rmsnorm_h4096 kernel.

    This function robustly handles both positional and keyword arguments,
    and passes them to the underlying implementation.

    Args:
        hidden_states (torch.Tensor): The input tensor of shape [batch_size, 4096]
                                      and dtype bfloat16.
        weight (torch.Tensor): The weight tensor of shape [4096] and dtype bfloat16.
    
    Returns:
        torch.Tensor: The normalized output tensor of the same shape and dtype as hidden_states.
    """
    try:
        # Use inspect to robustly bind arguments to the implementation function's signature.
        # This provides clear error messages for incorrect argument passing.
        sig = inspect.signature(_rmsnorm_h4096_impl)
        bound_args = sig.bind(*args, **kwargs)
        bound_args.apply_defaults()
    except TypeError as e:
        # Re-raise with a more user-friendly message.
        raise TypeError(f"Error binding arguments for rmsnorm_h4096: {e}") from e

    # Call the implementation with the correctly bound arguments.
    return _rmsnorm_h4096_impl(**bound_args.arguments)
scrolls · 175 lines total

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

Best evidence level for this revision: reported

JSON