Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonnicff0

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-nicff0?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 h512bf16 · [512] · batch_size=32
NVIDIA B200
6.18µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=64
NVIDIA B200
6.19µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=18
NVIDIA B200
6.20µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=7
NVIDIA B200
6.29µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=539
NVIDIA B200
6.35µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=1
NVIDIA B200
6.50µs
#4 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=11949
NVIDIA B200
13.5µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=14521
NVIDIA B200
14.5µs
#2 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:9e6abc101857a5b48b54bc94b2ddfae7697626d445ca6570f5819fe98f96489e
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 = 4num_warps=4 # A good default for a block size of 512.

Kernel source

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

# No hexadecimal float literals are used in this file.

@triton.jit
def rmsnorm_h512_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    stride_hidden_states_batch,
    stride_output_batch,
    HIDDEN_SIZE: tl.constexpr,
    EPS: tl.constexpr,
):
    """
    Triton kernel for RMSNorm.
    Each program instance processes a single row of the input tensor.
    """
    # The grid is 1D, so we can use `tl.program_id(0)` to get the row index.
    pid = tl.program_id(axis=0)

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

    # Load the entire row of hidden_states. Since HIDDEN_SIZE is 512,
    # a single program instance can handle it efficiently.
    offsets = tl.arange(0, HIDDEN_SIZE)
    
    # Load input data and promote to float32 for high-precision computation.
    # This is crucial for the stability of the variance calculation.
    x = tl.load(row_hidden_states_ptr + offsets, mask=offsets < HIDDEN_SIZE).to(tl.float32)

    # --- Start of RMSNorm computation ---
    
    # 1. Calculate the sum of squares.
    # This is a block-level reduction performed efficiently by `tl.sum`.
    sum_of_squares = tl.sum(x * x, axis=0)
    
    # 2. Compute the mean and add epsilon.
    mean_of_squares = sum_of_squares / HIDDEN_SIZE
    
    # 3. Calculate the reciprocal square root (rsqrt) for normalization.
    inv_rms = tl.rsqrt(mean_of_squares + EPS)
    
    # 4. Normalize the input row.
    # The `inv_rms` is a scalar, which is automatically broadcasted across the vector `x`.
    normalized_x = x * inv_rms
    
    # 5. Scale by the learnable weight parameter.
    # Load the weights and promote to float32 for the multiplication.
    w = tl.load(weight_ptr + offsets, mask=offsets < HIDDEN_SIZE).to(tl.float32)
    output_f32 = normalized_x * w

    # --- End of RMSNorm computation ---

    # Cast back to the original bfloat16 dtype and store the result.
    output = output_f32.to(tl.bfloat16)
    tl.store(row_output_ptr + offsets, output, mask=offsets < HIDDEN_SIZE)


def run(hidden_states: torch.Tensor, weight: torch.Tensor, *args, **kwargs) -> torch.Tensor:
    """
    Wrapper function for the RMSNorm Triton kernel.

    This function handles device management, input validation, kernel launching,
    and ensures the output tensor is on the same device as the input.

    Args:
        hidden_states (torch.Tensor): Input tensor of shape [batch_size, 512] and dtype bfloat16.
        weight (torch.Tensor): Weight tensor of shape [512] and dtype bfloat16.
        *args: Additional positional arguments (ignored).
        **kwargs: Additional keyword arguments (ignored).

    Returns:
        torch.Tensor: The output tensor of the same shape and dtype as `hidden_states`.
    """
    # --- Device Management & Input Validation ---
    if not torch.cuda.is_available():
        raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")

    # Validate tensor shapes and dtypes as per the specification.
    assert hidden_states.shape[1] == 512, f"Expected hidden_size=512, but got {hidden_states.shape[1]}"
    assert weight.shape == (512,), f"Expected weight shape (512,), but got {weight.shape}"
    assert hidden_states.dtype == torch.bfloat16, f"Expected hidden_states dtype bfloat16, but got {hidden_states.dtype}"
    assert weight.dtype == torch.bfloat16, f"Expected weight dtype bfloat16, but got {weight.dtype}"

    # Preserve the original device of the input tensor.
    original_device = hidden_states.device
    is_cpu_input = original_device.type == 'cpu'

    # Move tensors to GPU if they are not already there.
    if is_cpu_input:
        hidden_states_gpu = hidden_states.cuda()
        weight_gpu = weight.cuda()
    elif original_device.type != 'cuda':
        # Handle cases like 'mps' or other devices by moving data to cuda for computation
        hidden_states_gpu = hidden_states.to('cuda')
        weight_gpu = weight.to('cuda')
    else:
        hidden_states_gpu = hidden_states
        weight_gpu = weight

    # --- Kernel Launch ---
    batch_size, hidden_size = hidden_states_gpu.shape
    
    # Allocate the output tensor on the GPU.
    output_gpu = torch.empty_like(hidden_states_gpu)
    
    # Ensure tensors are contiguous for optimal memory access patterns.
    # Triton kernels often assume dense, C-style memory layout.
    hidden_states_gpu = hidden_states_gpu.contiguous()
    weight_gpu = weight_gpu.contiguous()
    output_gpu = output_gpu.contiguous()

    # The grid is 1D, with one program instance per row in the batch.
    grid = (batch_size,)
    
    # The epsilon value for numerical stability.
    EPS = 1e-6

    # Launch the Triton kernel.
    rmsnorm_h512_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,  # Pass as a compile-time constant for optimization.
        EPS=EPS,                  # Pass as a compile-time constant.
        num_warps=4               # A good default for a block size of 512.
    )

    # --- Result Handling ---
    # If the original input was on the CPU, move the result back to the CPU.
    if is_cpu_input:
        return output_gpu.to(original_device)
    
    return output_gpu
scrolls · 143 lines total

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

Best evidence level for this revision: reported

JSON