Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonxtl8hx

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-xtl8hx?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
Fused add RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
6.26µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
6.53µs
#2 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
6.76µs
#2 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
6.77µs
#3 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
31.0µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
34.2µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
38.8µ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:7ce758c012140b71baa84335500fbc0667e43d7add4fa05e47dbffaea84ef696
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.py179 lines
import torch
import triton
import triton.language as tl
import math

@triton.jit
def fused_add_rmsnorm_h2048_kernel(
    # Pointers to tensors
    hidden_states_ptr,
    residual_ptr,
    weight_ptr,
    output_ptr,
    # Stride variables for memory access
    stride_hidden_states_batch,
    stride_residual_batch,
    stride_output_batch,
    # Other parameters
    hidden_size,
    # Constants
    EPS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Triton kernel for fused Add + RMSNorm optimized for B200.
    - Each program instance processes one row of the input tensors.
    - It fuses the element-wise addition with the RMS normalization.
    - Computation is done in float32 for precision, while I/O is bfloat16.
    """
    # -----------------------------------------------------------
    # Map program ids `pid` to the batch dimension.
    # -----------------------------------------------------------
    # Each program instance handles one row.
    row_idx = tl.program_id(axis=0)

    # -----------------------------------------------------------
    # Pointers to the current row
    # -----------------------------------------------------------
    hidden_states_row_ptr = hidden_states_ptr + row_idx * stride_hidden_states_batch
    residual_row_ptr = residual_ptr + row_idx * stride_residual_batch
    output_row_ptr = output_ptr + row_idx * stride_output_batch

    # -----------------------------------------------------------
    # Load the row of data and compute the sum of squares
    # -----------------------------------------------------------
    # Create a block of offsets for the hidden dimension.
    # Since BLOCK_SIZE is expected to be equal to hidden_size, we load the whole row.
    offs = tl.arange(0, BLOCK_SIZE)
    
    # Load the input row vectors, casting to float32 for computation.
    hidden_states = tl.load(hidden_states_row_ptr + offs).to(tl.float32)
    residual = tl.load(residual_row_ptr + offs).to(tl.float32)

    # Fused add operation
    x = hidden_states + residual

    # Compute sum of squares for the variance calculation.
    # This is a scalar value for the row after the reduction.
    var = tl.sum(x * x, axis=0)
    
    # -----------------------------------------------------------
    # Compute RMS and apply normalization
    # -----------------------------------------------------------
    # Calculate inverse root mean square.
    rstd = tl.rsqrt(var / hidden_size + EPS)

    # Load the weight vector. It is broadcasted across all rows.
    weight = tl.load(weight_ptr + offs).to(tl.float32)

    # Normalize x and apply the learned scaling (weight).
    output_f32 = x * rstd * weight
    
    # -----------------------------------------------------------
    # Write the output
    # -----------------------------------------------------------
    # Cast back to the output dtype (bfloat16) and store.
    tl.store(output_row_ptr + offs, output_f32.to(tl.bfloat16))


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

    Handles device management, tensor validation, grid computation, and kernel launch.
    It moves tensors to the GPU, runs the kernel, and returns the result on the
    original device of the first input tensor.

    Args:
        hidden_states (torch.Tensor): The main input tensor of shape [batch_size, 2048] and dtype bfloat16.
        residual (torch.Tensor): The residual tensor to be added, with the same shape and dtype as hidden_states.
        weight (torch.Tensor): The scaling weights of shape [2048] and dtype bfloat16.

    Returns:
        torch.Tensor: The output tensor with the same shape and dtype as hidden_states.
    """
    # 1. Resolve and validate arguments from args and kwargs
    # This allows for flexible calling conventions (positional or keyword).
    arg_names = ['hidden_states', 'residual', 'weight']
    
    if args:
        if len(args) > len(arg_names):
            raise TypeError(f"run() takes at most {len(arg_names)} positional arguments but {len(args)} were given")
        for i, arg in enumerate(args):
            kwargs[arg_names[i]] = arg

    hidden_states = kwargs.get('hidden_states')
    residual = kwargs.get('residual')
    weight = kwargs.get('weight')

    if hidden_states is None or residual is None or weight is None:
        missing = [name for name in arg_names if name not in kwargs]
        raise TypeError(f"run() missing required arguments: {', '.join(missing)}")

    # 2. Device Management: determine target device and move tensors
    if not torch.cuda.is_available():
        raise RuntimeError("Triton kernel requires a CUDA-enabled GPU.")

    initial_devices = {
        'hidden_states': hidden_states.device,
        'residual': residual.device,
        'weight': weight.device
    }

    # Determine the target CUDA device. If any tensor is on CUDA, use that device.
    # Otherwise, default to the current CUDA device.
    target_device = None
    for tensor in [hidden_states, residual, weight]:
        if tensor.is_cuda:
            if target_device is None:
                target_device = tensor.device
            elif target_device != tensor.device:
                raise ValueError("All input tensors must be on the same CUDA device.")
    
    if target_device is None:
        target_device = torch.device("cuda")

    # Move all tensors to the target device for the kernel execution.
    hidden_states_gpu = hidden_states.to(target_device)
    residual_gpu = residual.to(target_device)
    weight_gpu = weight.to(target_device)

    # 3. Shape and DType validation on the device
    B, H = hidden_states_gpu.shape
    
    assert H == 2048, f"Expected hidden_size=2048, but got {H}"
    assert hidden_states_gpu.shape == residual_gpu.shape, "hidden_states and residual must have the same shape"
    assert weight_gpu.shape == (H,), f"Expected weight shape ({H},), but got {weight_gpu.shape}"
    assert hidden_states_gpu.ndim == 2, "Inputs must be 2D tensors"

    for name, tensor in [('hidden_states', hidden_states_gpu), ('residual', residual_gpu), ('weight', weight_gpu)]:
        if tensor.dtype != torch.bfloat16:
            raise TypeError(f"Expected {name} to be bfloat16, but got {tensor.dtype}")
    
    # 4. Allocate output tensor on the target device
    output = torch.empty_like(hidden_states_gpu)

    # 5. Set up grid and launch the kernel
    grid = (B,)
    
    # B200 benefits from more warps for large memory-bound operations.
    # For a large block size of 2048, 8 warps is a good choice.
    num_warps = 8
    
    fused_add_rmsnorm_h2048_kernel[grid](
        hidden_states_gpu,
        residual_gpu,
        weight_gpu,
        output,
        hidden_states_gpu.stride(0),
        residual_gpu.stride(0),
        output.stride(0),
        H,
        EPS=1e-6,
        BLOCK_SIZE=H,
        num_warps=num_warps
    )

    # 6. Move the result back to the original device of the primary input tensor
    output_device = initial_devices['hidden_states']
    return output.to(output_device)
scrolls · 179 lines total

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

Best evidence level for this revision: reported

JSON