gemini-2.5-pro / triton0owo26
gemini-2.5-pro_triton_0owo26 · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 175 lines ↓holds 5 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 175 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-0owo26?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
Show all 14 measurements ›Showing all 14 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c5d6dcefab0abc44d57188a1f74800b43548a83d568cb3033b4e32a30480650d
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 = 8
num_warps=8, # A good value for B200 to ensure high occupancyKernel source
main.py175 lines
import torch
import triton
import triton.language as tl
import math
import inspect
#
# Triton Kernel
#
@triton.jit
def rmsnorm_h4096_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
stride_hidden_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 of 4096.
This kernel is designed for B200 and performs the operation in a single fused pass
from the perspective of HBM. It reads hidden_states and weights once, and writes
the output once. The intermediate calculations (sum of squares, inv_rms) are
done in registers/SRAM.
Grid:
- The grid is 1D, with one program instance per batch element (row).
Memory Access:
- Each program loads an entire row of hidden_states (4096 elements).
- It also loads the entire weight vector (4096 elements).
- Accesses are contiguous and coalesced.
Computation:
- All intermediate computations are performed in float32 for precision.
- The final result is cast back to bfloat16 before storing.
"""
# Get the index for the current batch element.
batch_idx = tl.program_id(axis=0)
# --- Pointer Setup ---
# Pointers to the start of the current row for input and output tensors.
row_hidden_states_ptr = hidden_states_ptr + batch_idx * stride_hidden_batch
row_output_ptr = output_ptr + batch_idx * stride_output_batch
# Offsets for accessing the hidden dimension.
h_offsets = tl.arange(0, BLOCK_SIZE_H)
# --- Pass 1: Compute Inverse RMS ---
# Load the entire row of hidden_states for the current batch element.
# Since HIDDEN_SIZE is fixed at 4096, we load it in one go.
x_bf16 = tl.load(row_hidden_states_ptr + h_offsets)
x_fp32 = x_bf16.to(tl.float32)
# Calculate the sum of squares. tl.sum performs a parallel reduction
# within the thread block.
sum_sq = tl.sum(x_fp32 * x_fp32, axis=0)
# Calculate the mean and the inverse root mean square.
mean_sq = sum_sq / HIDDEN_SIZE
inv_rms = tl.rsqrt(mean_sq + EPS)
# --- Pass 2: Apply Normalization and Store ---
# Load the weight vector.
w_bf16 = tl.load(weight_ptr + h_offsets)
w_fp32 = w_bf16.to(tl.float32)
# Apply the RMS normalization and scale by the weights.
# The original x_fp32 values should still be in registers.
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(row_output_ptr + h_offsets, output_bf16)
#
# Wrapper Function
#
def _rmsnorm_h4096_impl(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""
Wrapper for the RMSNorm Triton kernel.
This function handles device management, shape/dtype validation,
grid computation, and kernel invocation.
"""
# --- Device Management & Validation ---
if not torch.cuda.is_available():
raise RuntimeError("This implementation requires a CUDA-enabled GPU.")
original_device = hidden_states.device
is_cpu_input = original_device.type == 'cpu'
# Move tensors to GPU if they are on CPU
if is_cpu_input:
hidden_states = hidden_states.cuda()
weight = weight.cuda()
# Validate shapes and dtypes
batch_size, hidden_size = hidden_states.shape
if hidden_size != 4096:
raise ValueError(f"This kernel is specialized for hidden_size=4096, but got {hidden_size}.")
if weight.shape != (hidden_size,):
raise ValueError(f"Weight shape mismatch. Expected ({hidden_size},) but got {weight.shape}.")
if hidden_states.dtype != torch.bfloat16:
raise TypeError(f"Input 'hidden_states' must be of dtype torch.bfloat16, but got {hidden_states.dtype}.")
if weight.dtype != torch.bfloat16:
raise TypeError(f"Input 'weight' must be of dtype torch.bfloat16, but got {weight.dtype}.")
# --- Kernel Launch ---
# Allocate the output tensor on the same GPU device.
output = torch.empty_like(hidden_states)
# Define the grid for the kernel launch. One program per batch element.
grid = (batch_size,)
# Constants for the kernel.
HIDDEN_SIZE = 4096
EPS = 1e-5
# Launch the Triton kernel.
rmsnorm_h4096_kernel[grid](
hidden_states,
weight,
output,
hidden_states.stride(0),
output.stride(0),
HIDDEN_SIZE=HIDDEN_SIZE,
EPS=EPS,
BLOCK_SIZE_H=HIDDEN_SIZE, # Process the entire row in one block
num_warps=8, # A good value for B200 to ensure high occupancy
)
# --- Final Device Management ---
# If the original input was on CPU, move the result back to CPU.
if is_cpu_input:
output = output.to(original_device)
return output
#
# Public Entry Point
#
def run(*args, **kwargs):
"""
Public entry point for the rmsnorm_h4096 kernel.
This function robustly handles both positional and keyword arguments,
and passes them to the underlying implementation.
Args:
hidden_states (torch.Tensor): The input tensor of shape [batch_size, 4096]
and dtype bfloat16.
weight (torch.Tensor): The weight tensor of shape [4096] and dtype bfloat16.
Returns:
torch.Tensor: The normalized output tensor of the same shape and dtype as hidden_states.
"""
try:
# Use inspect to robustly bind arguments to the implementation function's signature.
# This provides clear error messages for incorrect argument passing.
sig = inspect.signature(_rmsnorm_h4096_impl)
bound_args = sig.bind(*args, **kwargs)
bound_args.apply_defaults()
except TypeError as e:
# Re-raise with a more user-friendly message.
raise TypeError(f"Error binding arguments for rmsnorm_h4096: {e}") from e
# Call the implementation with the correctly bound arguments.
return _rmsnorm_h4096_impl(**bound_args.arguments)
scrolls · 175 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON