Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonuc5ufi

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

7 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
6.17µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
6.18µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
6.21µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
6.22µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
6.27µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
25.7µs
#1 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
31.1µ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:ea78c97c1533ea7cfda3af7b24a85e403402852045b3705de876a01da3013a4c
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.py181 lines
import torch
import triton
import triton.language as tl

@triton.jit
def rmsnorm_h2048_kernel(
    # Pointers to tensors
    X_ptr,
    Weight_ptr,
    Output_ptr,
    # Stride to move to the next row
    stride_x_batch,
    stride_output_batch,
    # Constants
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
):
    """
    Triton kernel for RMS Normalization.
    This kernel is optimized for hidden_size=2048 on B200 GPUs.
    Each program instance handles a single row of the input tensor.
    """
    # --- Grid and Block Configuration ---
    # Each program instance processes one row of the input tensor.
    row_idx = tl.program_id(0)

    # --- Pointer Setup ---
    # Pointers to the start of the current row for input and output.
    x_row_ptr = X_ptr + row_idx * stride_x_batch
    output_row_ptr = Output_ptr + row_idx * stride_output_batch

    # Pointers to the columns of the current row.
    col_offsets = tl.arange(0, HIDDEN_SIZE)
    x_ptrs = x_row_ptr + col_offsets
    weight_ptrs = Weight_ptr + col_offsets
    output_ptrs = output_row_ptr + col_offsets

    # --- Computation ---
    # Load the input row and weights.
    # For HIDDEN_SIZE=2048, the entire row and weight vector fit into SRAM,
    # enabling a single-pass algorithm.
    x_bf16 = tl.load(x_ptrs)
    w_bf16 = tl.load(weight_ptrs)

    # Upcast to float32 for high-precision calculation of variance.
    x_fp32 = x_bf16.to(tl.float32)
    
    # Calculate the sum of squares.
    # This is a block-level reduction handled efficiently by Triton.
    sum_of_squares = tl.sum(x_fp32 * x_fp32, axis=0)
    
    # Calculate variance and the inverse root mean square.
    variance = sum_of_squares / HIDDEN_SIZE
    inv_rms = tl.rsqrt(variance + EPS)

    # Normalize the input and apply the learned scaling factor (weight).
    # Reuse the loaded `x_fp32` from registers/SRAM.
    w_fp32 = w_bf16.to(tl.float32)
    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(output_ptrs, output_bf16)


def rmsnorm_h2048(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
    """
    Host-side wrapper for the RMSNorm Triton kernel.

    This function handles device management, input validation, kernel launching,
    and moving the final result back to the original device.
    
    Args:
        hidden_states (torch.Tensor): Input tensor of shape [batch_size, 2048] and dtype bfloat16.
        weight (torch.Tensor): Weight tensor of shape [2048] and dtype bfloat16.
        
    Returns:
        torch.Tensor: The normalized output tensor of the same shape and dtype as the input.
    """
    # --- Device Management and Input Validation ---
    original_device = hidden_states.device
    
    # Triton requires a CUDA-enabled GPU.
    if not torch.cuda.is_available():
        if original_device.type == 'cuda':
            raise RuntimeError("CUDA is not available, but input tensor is on a CUDA device.")
        raise RuntimeError("Triton requires a CUDA-enabled GPU, but none was found.")
    
    # Move tensors to the GPU if they are on the CPU.
    device = torch.device("cuda")
    hidden_states_gpu = hidden_states.to(device) if hidden_states.device.type == 'cpu' else hidden_states
    weight_gpu = weight.to(device) if weight.device.type == 'cpu' else weight

    # Ensure all tensors are on the same CUDA device.
    if weight_gpu.device != hidden_states_gpu.device:
        weight_gpu = weight_gpu.to(hidden_states_gpu.device)
        
    # --- Shape and DType Validation ---
    batch_size, hidden_size = hidden_states_gpu.shape
    
    if hidden_size != 2048:
        raise ValueError(f"Expected hidden_size=2048, but got {hidden_size}")
    if weight_gpu.shape != (2048,):
        raise ValueError(f"Expected weight shape=(2048,), but got {weight_gpu.shape}")
    if hidden_states_gpu.dtype != torch.bfloat16:
        raise TypeError(f"Expected hidden_states dtype=torch.bfloat16, but got {hidden_states_gpu.dtype}")
    if weight_gpu.dtype != torch.bfloat16:
        raise TypeError(f"Expected weight dtype=torch.bfloat16, but got {weight_gpu.dtype}")
    
    # Ensure inputs are contiguous for safe memory access.
    if not hidden_states_gpu.is_contiguous():
        hidden_states_gpu = hidden_states_gpu.contiguous()
    if not weight_gpu.is_contiguous():
        weight_gpu = weight_gpu.contiguous()

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

    # The grid is 1D, with one program instance per row in the batch.
    grid = (batch_size,)
    
    # Constants for the kernel.
    EPS = 1e-6
    
    # For modern GPUs like B200, 8 warps is a good choice for a 2048-element reduction.
    # This gives 256 threads, each handling 8 elements, balancing parallelism and per-thread work.
    num_warps = 8

    # Call the Triton kernel.
    rmsnorm_h2048_kernel[grid](
        hidden_states_gpu,
        weight_gpu,
        output,
        hidden_states_gpu.stride(0),
        output.stride(0),
        HIDDEN_SIZE=hidden_size,
        EPS=EPS,
        num_warps=num_warps,
    )

    # --- Post-computation Device Management ---
    # Move the output back to the original device of the input tensor.
    if output.device != original_device:
        output = output.to(original_device)
        
    return output


def run(*args, **kwargs) -> torch.Tensor:
    """
    Public entry point for the rmsnorm_h2048 operation.
    
    This function handles generic `*args` and `**kwargs` inputs, dispatches
    to the main `rmsnorm_h2048` implementation, and ensures compatibility
    with various calling conventions.
    
    Args:
        *args: Positional arguments. Expected: `hidden_states`, `weight`.
        **kwargs: Keyword arguments. Expected: `hidden_states`, `weight`.
        
    Returns:
        torch.Tensor: The result of the RMS normalization, with the same shape,
                      dtype, and device as the `hidden_states` input.
    """
    # Extract arguments by name from kwargs first, then fall back to args.
    hidden_states = kwargs.get('hidden_states')
    if hidden_states is None:
        if len(args) > 0:
            hidden_states = args[0]
        else:
            raise TypeError("run() missing 1 required positional argument: 'hidden_states'")

    weight = kwargs.get('weight')
    if weight is None:
        if len(args) > 1:
            weight = args[1]
        else:
            raise TypeError("run() missing 1 required positional argument: 'weight'")
            
    return rmsnorm_h2048(hidden_states, weight)
scrolls · 181 lines total

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

Best evidence level for this revision: reported

JSON