Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_triton_xndzsl

gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-xndzsl?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 h1536bf16 · [1536] · batch_size=64
NVIDIA B200
6.52µs
#2 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=32
NVIDIA B200
7.39µs
#3 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=18
NVIDIA B200
7.45µs
#3 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=1
NVIDIA B200
7.52µs
#3 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=539
NVIDIA B200
8.31µs
#3 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=7
NVIDIA B200
8.34µs
#3 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=11949
NVIDIA B200
18.5µs
#1 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=14521
NVIDIA B200
20.6µ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:c3b7d3dd9cafda3f27ff4686f55cd6dea9203892dfea2d09bbdb56ba1c11dd67
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({}, num_warps=4),

Kernel source

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

# Reference implementation for fallback and verification
def _reference_run(hidden_states, weight):
    """
    Reference PyTorch implementation for RMSNorm.
    """
    batch_size, hidden_size = hidden_states.shape
    assert hidden_size == 1536

    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.autotune(
    configs=[
        triton.Config({}, num_warps=4),
        triton.Config({}, num_warps=8),
        triton.Config({}, num_warps=16),
    ],
    key=['HIDDEN_SIZE'],
)
@triton.jit
def _rmsnorm_h1536_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    stride_hidden_states_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.
    
    This kernel is fused to perform the normalization in a single pass over the data,
    minimizing global memory access. It processes one row per program instance.
    
    1. It computes the sum of squares for a row.
    2. It calculates the inverse root mean square (inv_rms).
    3. It normalizes the input row with inv_rms, scales it by the weight vector,
       and writes the result to the output.
    
    The use of tl.float32 for intermediate calculations (sum of squares) ensures
    numerical accuracy.
    """
    # Grid is 1D, so each program instance processes one row.
    pid_row = tl.program_id(0)

    # Pointers to the current row for inputs and output.
    row_hidden_states_ptr = hidden_states_ptr + pid_row * stride_hidden_states_batch
    row_output_ptr = output_ptr + pid_row * stride_output_batch

    # --- 1. Compute mean of squares ---
    # Create a block of offsets for the hidden dimension.
    offs_h = tl.arange(0, BLOCK_SIZE_H)
    mask_h = offs_h < HIDDEN_SIZE

    # Load the row of hidden_states. Use masking to handle cases where
    # BLOCK_SIZE_H > HIDDEN_SIZE. `other=0.0` ensures that padding
    # doesn't affect the sum of squares.
    x_block = tl.load(row_hidden_states_ptr + offs_h, mask=mask_h, other=0.0)
    
    # Promote to float32 for high-precision reduction.
    x_f32 = x_block.to(tl.float32)
    
    # Calculate sum of squares.
    sum_sq = tl.sum(x_f32 * x_f32, axis=0)
    
    # Calculate mean of squares. We divide by the actual HIDDEN_SIZE.
    mean_sq = sum_sq / HIDDEN_SIZE
    
    # --- 2. Compute inverse root mean square ---
    # Add epsilon for numerical stability and compute rsqrt.
    inv_rms = tl.rsqrt(mean_sq + EPS)

    # --- 3. Normalize, scale, and store ---
    # Load the corresponding weights.
    weight_block = tl.load(weight_ptr + offs_h, mask=mask_h)
    
    # Perform normalization and scaling.
    # We reuse x_f32, which is already in registers from the initial load.
    output_f32 = x_f32 * inv_rms * weight_block.to(tl.float32)

    # Convert back to the output dtype (bfloat16) and store.
    tl.store(row_output_ptr + offs_h, output_f32.to(tl.bfloat16), mask=mask_h)


def run(*args, **kwargs):
    """
    Wrapper function for the RMSNorm Triton kernel.

    Handles device management, argument parsing, and kernel launching.
    It ensures that tensors are on the correct device (CUDA) for kernel
    execution and that the output is moved back to the original device
    of the input tensors.

    Args:
        hidden_states (torch.Tensor): The input tensor of shape [batch_size, 1536]
                                      and dtype bfloat16.
        weight (torch.Tensor): The weight tensor of shape [1536] and dtype bfloat16.
    
    Returns:
        torch.Tensor: The normalized and scaled output tensor with the same shape
                      and dtype as hidden_states.
    """
    # --- Argument Parsing ---
    hidden_states = kwargs.get('hidden_states')
    weight = kwargs.get('weight')
    
    if len(args) > 0:
        if hidden_states is not None:
            raise TypeError("run() got multiple values for argument 'hidden_states'")
        hidden_states = args[0]
    if len(args) > 1:
        if weight is not None:
            raise TypeError("run() got multiple values for argument 'weight'")
        weight = args[1]
    if len(args) > 2:
        raise TypeError(f"run() takes 2 positional arguments but {len(args)} were given")
    
    if hidden_states is None or weight is None:
        raise TypeError("run() missing required arguments: 'hidden_states' or 'weight'")

    # --- Input Validation ---
    assert isinstance(hidden_states, torch.Tensor), "hidden_states must be a torch.Tensor"
    assert isinstance(weight, torch.Tensor), "weight must be a torch.Tensor"
    assert hidden_states.shape[1] == 1536, "hidden_size must be 1536"
    assert hidden_states.shape[1] == weight.shape[0], "hidden_states and weight dimensions must match"
    assert hidden_states.ndim == 2, "hidden_states must be a 2D tensor"
    assert weight.ndim == 1, "weight must be a 1D tensor"
    assert hidden_states.dtype == torch.bfloat16, "hidden_states dtype must be bfloat16"
    assert weight.dtype == torch.bfloat16, "weight dtype must be bfloat16"

    # --- Device Management ---
    original_device = hidden_states.device
    
    if not torch.cuda.is_available():
        if original_device.type == 'cuda':
             raise RuntimeError("CUDA is not available, but input tensors are on a CUDA device.")
        # As Triton is unavailable, fall back to a native PyTorch implementation on CPU.
        print("Warning: Triton requires a CUDA-enabled GPU. Falling back to reference implementation on CPU.")
        return _reference_run(hidden_states, weight)

    device = torch.device("cuda")
    hidden_states_gpu = hidden_states.to(device)
    weight_gpu = weight.to(device)

    # --- Allocate Output Tensor ---
    output_gpu = torch.empty_like(hidden_states_gpu)
    
    # --- Kernel Launch ---
    batch_size, hidden_size = hidden_states_gpu.shape
    
    # Use a block size that is the next power of 2 to allow for efficient hardware utilization
    BLOCK_SIZE_H = triton.next_power_of_2(hidden_size)

    grid = (batch_size, )
    
    _rmsnorm_h1536_kernel[grid](
        hidden_states_ptr=hidden_states_gpu,
        weight_ptr=weight_gpu,
        output_ptr=output_gpu,
        stride_hidden_states_batch=hidden_states_gpu.stride(0),
        stride_output_batch=output_gpu.stride(0),
        HIDDEN_SIZE=hidden_size,
        EPS=1e-6,
        BLOCK_SIZE_H=BLOCK_SIZE_H,
    )
    
    # --- Move Result to Original Device ---
    return output_gpu.to(original_device)
scrolls · 180 lines total

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

Best evidence level for this revision: reported

JSON