Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_triton_jaedrq

gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

14 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h128bf16 · [128] · batch_size=4
NVIDIA B200
6.18µs
#3= of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=32
NVIDIA B200
6.18µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=256
NVIDIA B200
6.19µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=136
NVIDIA B200
6.19µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=316
NVIDIA B200
6.20µs
#4 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=24
NVIDIA B200
6.20µs
#4 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=1088
NVIDIA B200
6.20µs
#1 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=192
NVIDIA B200
6.21µs
#4 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2048
NVIDIA B200
7.29µs
#1 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2528
NVIDIA B200
8.12µs
#3 of 9
2025-10-16
Show all 14 measurements ›
RMSNorm h128bf16 · [128] · batch_size=49532
NVIDIA B200
32.5µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=65016
NVIDIA B200
39.6µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=396256
NVIDIA B200
211.0µs
#5 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=520128
NVIDIA B200
274.6µs
#5 of 9
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:8110ee2a517d85eeab6a1ac5b612cca680305217c75945cae1c18c770ce3e834
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,

Kernel source

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

#
# Reference Implementation:
# import torch
#
# @torch.no_grad()
# def run(hidden_states, weight):
#     batch_size, hidden_size = hidden_states.shape
#     # Check constants
#     assert hidden_size == 128
#
#     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_h128_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    hidden_size,
    stride_b,
    EPS: tl.constexpr,
    BLOCK_SIZE_H: tl.constexpr,
):
    """
    Triton kernel for RMS Normalization with a fixed hidden_size of 128.
    Each program in the grid processes one row from the batch.
    """
    # Get the row index for the current program
    pid_b = tl.program_id(axis=0)

    # Create pointers to the start of the current row for inputs and output
    row_x_ptr = hidden_states_ptr + pid_b * stride_b
    row_y_ptr = output_ptr + pid_b * stride_b

    # Create a range of offsets for the hidden dimension
    # Since BLOCK_SIZE_H is fixed to hidden_size (128), we load the entire row
    offsets_h = tl.arange(0, BLOCK_SIZE_H)
    
    # Load the full row of hidden_states and the full weight vector
    # No mask is needed as hidden_size == BLOCK_SIZE_H
    x = tl.load(row_x_ptr + offsets_h)
    w = tl.load(weight_ptr + offsets_h)

    # --- Computation is performed in float32 for precision ---
    x_fp32 = x.to(tl.float32)
    w_fp32 = w.to(tl.float32)

    # 1. Square the elements
    x_sq = x_fp32 * x_fp32

    # 2. Compute the mean of the squares (reduction)
    # tl.sum performs an efficient reduction over the block of 128 elements
    var = tl.sum(x_sq, axis=0) / hidden_size

    # 3. Compute the inverse root mean square
    inv_rms = tl.rsqrt(var + EPS)

    # 4. Normalize the hidden states and apply the learned scaling factor (weight)
    y = x_fp32 * inv_rms * w_fp32

    # --- Cast back to bfloat16 and store the result ---
    y_bf16 = y.to(tl.bfloat16)
    tl.store(row_y_ptr + offsets_h, y_bf16)


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

    Handles device management, tensor validation, and kernel launching.
    It preserves the device of the input tensors for the output.

    Args:
        hidden_states (torch.Tensor): Input tensor of shape [batch_size, 128]
                                      and dtype bfloat16.
        weight (torch.Tensor): Weight tensor of shape [128] and dtype bfloat16.
    
    Returns:
        torch.Tensor: The normalized output tensor of the same shape and dtype
                      as hidden_states.
    """
    # 1. Parse arguments
    if args:
        if len(args) != 2:
            raise ValueError(f"Expected 2 positional arguments, but got {len(args)}")
        hidden_states, weight = args
    else:
        hidden_states = kwargs.get('hidden_states')
        weight = kwargs.get('weight')
        if hidden_states is None or weight is None:
            raise ValueError("Missing required keyword arguments: 'hidden_states' and/or 'weight'")

    # 2. Validate tensor properties
    if hidden_states.dim() != 2 or hidden_states.shape[1] != 128:
        raise ValueError(f"Expected hidden_states to have shape [batch_size, 128], but got {hidden_states.shape}")
    if weight.dim() != 1 or weight.shape[0] != 128:
        raise ValueError(f"Expected weight to have shape [128], but got {weight.shape}")
    if hidden_states.dtype != torch.bfloat16:
        raise TypeError(f"Expected hidden_states to have dtype torch.bfloat16, but got {hidden_states.dtype}")
    if weight.dtype != torch.bfloat16:
        raise TypeError(f"Expected weight to have dtype torch.bfloat16, but got {weight.dtype}")

    # 3. Device management
    original_device = hidden_states.device
    is_cpu_input = original_device.type == 'cpu'
    
    # If inputs are on CPU, they must be moved to a GPU to run the Triton kernel.
    if is_cpu_input:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
        target_device = 'cuda'
        hidden_states = hidden_states.to(target_device)
        weight = weight.to(target_device)
    # If inputs are already on a GPU, ensure they are on the same device.
    else:
        if hidden_states.device != weight.device:
            raise ValueError(f"hidden_states and weight must be on the same device, "
                             f"but got {hidden_states.device} and {weight.device}")
        target_device = hidden_states.device

    # 4. Prepare for kernel launch
    batch_size, hidden_size = hidden_states.shape
    
    # Allocate the output tensor on the target GPU device
    output = torch.empty_like(hidden_states)

    # The grid is 1D, with one program per row in the batch.
    grid = (batch_size,)

    # Constants for the kernel
    EPS = 1e-6
    BLOCK_SIZE_H = 128

    # 5. Launch the Triton kernel
    # num_warps=4 is a robust choice for a block size of 128 on modern GPUs like B200.
    rmsnorm_h128_kernel[grid](
        hidden_states,
        weight,
        output,
        batch_size,
        hidden_size,
        hidden_states.stride(0),
        EPS=EPS,
        BLOCK_SIZE_H=BLOCK_SIZE_H,
        num_warps=4,
    )

    # 6. Move the result back to the original device if necessary
    if is_cpu_input:
        output = output.to(original_device)

    return output
scrolls · 163 lines total

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

Best evidence level for this revision: reported

JSON