gemini-2.5-pro / tritonuc5ufi
gemini-2.5-pro_triton_uc5ufi · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 181 lines ↓holds 7 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 181 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-uc5ufi?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
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ea78c97c1533ea7cfda3af7b24a85e403402852045b3705de876a01da3013a4c
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 = 8Kernel source
main.py181 lines
import torch
import triton
import triton.language as tl
@triton.jit
def rmsnorm_h2048_kernel(
# Pointers to tensors
X_ptr,
Weight_ptr,
Output_ptr,
# Stride to move to the next row
stride_x_batch,
stride_output_batch,
# Constants
HIDDEN_SIZE: tl.constexpr,
EPS: tl.constexpr,
):
"""
Triton kernel for RMS Normalization.
This kernel is optimized for hidden_size=2048 on B200 GPUs.
Each program instance handles a single row of the input tensor.
"""
# --- Grid and Block Configuration ---
# Each program instance processes one row of the input tensor.
row_idx = tl.program_id(0)
# --- Pointer Setup ---
# Pointers to the start of the current row for input and output.
x_row_ptr = X_ptr + row_idx * stride_x_batch
output_row_ptr = Output_ptr + row_idx * stride_output_batch
# Pointers to the columns of the current row.
col_offsets = tl.arange(0, HIDDEN_SIZE)
x_ptrs = x_row_ptr + col_offsets
weight_ptrs = Weight_ptr + col_offsets
output_ptrs = output_row_ptr + col_offsets
# --- Computation ---
# Load the input row and weights.
# For HIDDEN_SIZE=2048, the entire row and weight vector fit into SRAM,
# enabling a single-pass algorithm.
x_bf16 = tl.load(x_ptrs)
w_bf16 = tl.load(weight_ptrs)
# Upcast to float32 for high-precision calculation of variance.
x_fp32 = x_bf16.to(tl.float32)
# Calculate the sum of squares.
# This is a block-level reduction handled efficiently by Triton.
sum_of_squares = tl.sum(x_fp32 * x_fp32, axis=0)
# Calculate variance and the inverse root mean square.
variance = sum_of_squares / HIDDEN_SIZE
inv_rms = tl.rsqrt(variance + EPS)
# Normalize the input and apply the learned scaling factor (weight).
# Reuse the loaded `x_fp32` from registers/SRAM.
w_fp32 = w_bf16.to(tl.float32)
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(output_ptrs, output_bf16)
def rmsnorm_h2048(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""
Host-side wrapper for the RMSNorm Triton kernel.
This function handles device management, input validation, kernel launching,
and moving the final result back to the original device.
Args:
hidden_states (torch.Tensor): Input tensor of shape [batch_size, 2048] and dtype bfloat16.
weight (torch.Tensor): Weight tensor of shape [2048] and dtype bfloat16.
Returns:
torch.Tensor: The normalized output tensor of the same shape and dtype as the input.
"""
# --- Device Management and Input Validation ---
original_device = hidden_states.device
# Triton requires a CUDA-enabled GPU.
if not torch.cuda.is_available():
if original_device.type == 'cuda':
raise RuntimeError("CUDA is not available, but input tensor is on a CUDA device.")
raise RuntimeError("Triton requires a CUDA-enabled GPU, but none was found.")
# Move tensors to the GPU if they are on the CPU.
device = torch.device("cuda")
hidden_states_gpu = hidden_states.to(device) if hidden_states.device.type == 'cpu' else hidden_states
weight_gpu = weight.to(device) if weight.device.type == 'cpu' else weight
# Ensure all tensors are on the same CUDA device.
if weight_gpu.device != hidden_states_gpu.device:
weight_gpu = weight_gpu.to(hidden_states_gpu.device)
# --- Shape and DType Validation ---
batch_size, hidden_size = hidden_states_gpu.shape
if hidden_size != 2048:
raise ValueError(f"Expected hidden_size=2048, but got {hidden_size}")
if weight_gpu.shape != (2048,):
raise ValueError(f"Expected weight shape=(2048,), but got {weight_gpu.shape}")
if hidden_states_gpu.dtype != torch.bfloat16:
raise TypeError(f"Expected hidden_states dtype=torch.bfloat16, but got {hidden_states_gpu.dtype}")
if weight_gpu.dtype != torch.bfloat16:
raise TypeError(f"Expected weight dtype=torch.bfloat16, but got {weight_gpu.dtype}")
# Ensure inputs are contiguous for safe memory access.
if not hidden_states_gpu.is_contiguous():
hidden_states_gpu = hidden_states_gpu.contiguous()
if not weight_gpu.is_contiguous():
weight_gpu = weight_gpu.contiguous()
# --- Kernel Launch ---
# Allocate the output tensor on the GPU.
output = torch.empty_like(hidden_states_gpu)
# The grid is 1D, with one program instance per row in the batch.
grid = (batch_size,)
# Constants for the kernel.
EPS = 1e-6
# For modern GPUs like B200, 8 warps is a good choice for a 2048-element reduction.
# This gives 256 threads, each handling 8 elements, balancing parallelism and per-thread work.
num_warps = 8
# Call the Triton kernel.
rmsnorm_h2048_kernel[grid](
hidden_states_gpu,
weight_gpu,
output,
hidden_states_gpu.stride(0),
output.stride(0),
HIDDEN_SIZE=hidden_size,
EPS=EPS,
num_warps=num_warps,
)
# --- Post-computation Device Management ---
# Move the output back to the original device of the input tensor.
if output.device != original_device:
output = output.to(original_device)
return output
def run(*args, **kwargs) -> torch.Tensor:
"""
Public entry point for the rmsnorm_h2048 operation.
This function handles generic `*args` and `**kwargs` inputs, dispatches
to the main `rmsnorm_h2048` implementation, and ensures compatibility
with various calling conventions.
Args:
*args: Positional arguments. Expected: `hidden_states`, `weight`.
**kwargs: Keyword arguments. Expected: `hidden_states`, `weight`.
Returns:
torch.Tensor: The result of the RMS normalization, with the same shape,
dtype, and device as the `hidden_states` input.
"""
# Extract arguments by name from kwargs first, then fall back to args.
hidden_states = kwargs.get('hidden_states')
if hidden_states is None:
if len(args) > 0:
hidden_states = args[0]
else:
raise TypeError("run() missing 1 required positional argument: 'hidden_states'")
weight = kwargs.get('weight')
if weight is None:
if len(args) > 1:
weight = args[1]
else:
raise TypeError("run() missing 1 required positional argument: 'weight'")
return rmsnorm_h2048(hidden_states, weight)scrolls · 181 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON