Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton5h17k3

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

8 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
8.20µs
#3 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
8.21µs
#3 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=64
NVIDIA B200
8.23µs
#3 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=18
NVIDIA B200
8.25µs
#3 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
8.31µs
#3 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
10.2µs
#1 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
63.2µs
#1 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
73.7µs
#1 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:f6b8f61700fc285ee9b0aa0edcb2dbe9b3b78917612e7e0eec92f02b2a4f1159
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,

Kernel source

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

# Reference implementation for mathematical specification verification
# @torch.no_grad()
# def reference_run(hidden_states, weight):
#     batch_size, hidden_size = hidden_states.shape
#     # Check constants
#     assert hidden_size == 7168
#
#     EPS = 1e-6
#
#     x = hidden_states.to(torch.float32)
#     inv_rms = torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + EPS)
#     y = (x * inv_rms) * weight.to(torch.float32)
#     return y.to(hidden_states.dtype)

@triton.jit
def _rmsnorm_kernel(
    # Pointers to tensors
    x_ptr,
    weight_ptr,
    output_ptr,
    # Stride to move to the next row
    stride_x_batch,
    stride_out_batch,
    # Matrix dimensions
    hidden_size,
    # Kernel constants
    EPS: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
):
    """
    Triton kernel for RMS Normalization.
    This kernel is optimized for a fixed hidden_size and targets B200 GPUs.

    Grid: 1D, with each program instance processing one row (one item in the batch).
    """
    # Each program instance processes a single row.
    pid_batch = tl.program_id(axis=0)

    # Pointers to the current row for input and output.
    row_x_ptr = x_ptr + pid_batch * stride_x_batch
    row_output_ptr = output_ptr + pid_batch * stride_out_batch

    # B200 Optimization: Use a large block size to process the entire row in a single,
    # vectorized operation. This maximizes memory bandwidth utilization.
    # `BLOCK_SIZE_N` is configured to be the next power of 2 of `hidden_size`.
    offsets_n = tl.arange(0, BLOCK_SIZE_N)
    mask_n = offsets_n < hidden_size

    # --- Pass 1: Compute sum of squares and inv_rms ---

    # Load the entire row into registers (SRAM).
    # Convert to float32 for high-precision accumulation to avoid overflow/underflow.
    x = tl.load(row_x_ptr + offsets_n, mask=mask_n, other=0.0).to(tl.float32)

    # Compute the sum of squares. tl.sum performs an efficient reduction.
    sum_sq = tl.sum(x * x, axis=0)
    
    # Calculate variance and inverse root mean square.
    var = sum_sq / hidden_size
    inv_rms = tl.rsqrt(var + EPS)

    # --- Pass 2: Normalize, scale, and store ---
    
    # This pass is fused and operates on data held in registers.
    # Load the corresponding weights.
    w = tl.load(weight_ptr + offsets_n, mask=mask_n, other=0.0).to(tl.float32)

    # Apply the normalization and scaling.
    output_val = x * inv_rms * w

    # Convert back to the target dtype (bfloat16) and store the result.
    tl.store(row_output_ptr + offsets_n, output_val.to(tl.bfloat16), mask=mask_n)


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

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

    Returns:
        torch.Tensor: The normalized and scaled output tensor.
    """
    # Input validation
    if hidden_states.shape[1] != 7168:
        raise ValueError(f"Expected hidden_size=7168, but got {hidden_states.shape[1]}")
    if hidden_states.dtype != torch.bfloat16:
        raise TypeError(f"Expected hidden_states dtype bfloat16, but got {hidden_states.dtype}")
    if weight.shape != (7168,):
        raise ValueError(f"Expected weight shape (7168,), but got {weight.shape}")
    if weight.dtype != torch.bfloat16:
        raise TypeError(f"Expected weight dtype bfloat16, but got {weight.dtype}")

    # Kernel parameters
    batch_size, hidden_size = hidden_states.shape
    
    # Allocate output tensor
    output = torch.empty_like(hidden_states)

    # Grid definition: one program per row
    grid = (batch_size,)

    # B200 Optimization: Choose a block size that covers the entire row dimension.
    # This allows for full vectorization and avoids loop overhead within the kernel.
    BLOCK_SIZE_N = triton.next_power_of_2(hidden_size)

    # Kernel launch
    _rmsnorm_kernel[grid](
        x_ptr=hidden_states,
        weight_ptr=weight,
        output_ptr=output,
        stride_x_batch=hidden_states.stride(0),
        stride_out_batch=output.stride(0),
        hidden_size=hidden_size,
        EPS=1e-6,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        # B200 Optimization: Use a higher number of warps to hide memory latency,
        # which is crucial for memory-bound operations like this.
        num_warps=8,
    )

    return output


def run(*args, **kwargs):
    """
    Public entry point for the RMSNorm operation.
    Handles device management and calls the Triton kernel implementation.
    """
    # 1. Parse arguments
    if args:
        hidden_states, weight = args
    elif kwargs:
        hidden_states = kwargs.get('hidden_states')
        weight = kwargs.get('weight')
    else:
        raise ValueError("Missing required arguments 'hidden_states' and 'weight'")

    if hidden_states is None or weight is None:
        raise ValueError("Both 'hidden_states' and 'weight' must be provided")

    # 2. Device Management: Setup
    original_device = hidden_states.device
    
    if not torch.cuda.is_available():
        raise RuntimeError("Triton requires CUDA, but torch.cuda.is_available() is False.")
    
    target_device = torch.device("cuda")

    # 3. Move tensors to GPU if they aren't already
    inputs_on_gpu = True
    if hidden_states.device != target_device:
        hidden_states = hidden_states.to(target_device)
        inputs_on_gpu = False
    if weight.device != target_device:
        weight = weight.to(target_device)
        inputs_on_gpu = False

    # 4. Execute the kernel
    output = rmsnorm_h7168(hidden_states, weight)

    # 5. Device Management: Move result back to the original device
    if original_device != target_device:
        output = output.to(original_device)
        
    return output
scrolls · 174 lines total

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

Best evidence level for this revision: reported

JSON