gemini-2.5-pro_triton_xndzsl
gemini-2.5-pro · triton · Apache-2.0
Kernel source · 180 lines ↓holds 2 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 180 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-xndzsl?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
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c3b7d3dd9cafda3f27ff4686f55cd6dea9203892dfea2d09bbdb56ba1c11dd67
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.
autotune
@triton.autotune(num-warps = 4
triton.Config({}, num_warps=4),Kernel source
main.py180 lines
import torch
import triton
import triton.language as tl
import math
# Reference implementation for fallback and verification
def _reference_run(hidden_states, weight):
"""
Reference PyTorch implementation for RMSNorm.
"""
batch_size, hidden_size = hidden_states.shape
assert hidden_size == 1536
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.autotune(
configs=[
triton.Config({}, num_warps=4),
triton.Config({}, num_warps=8),
triton.Config({}, num_warps=16),
],
key=['HIDDEN_SIZE'],
)
@triton.jit
def _rmsnorm_h1536_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
stride_hidden_states_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.
This kernel is fused to perform the normalization in a single pass over the data,
minimizing global memory access. It processes one row per program instance.
1. It computes the sum of squares for a row.
2. It calculates the inverse root mean square (inv_rms).
3. It normalizes the input row with inv_rms, scales it by the weight vector,
and writes the result to the output.
The use of tl.float32 for intermediate calculations (sum of squares) ensures
numerical accuracy.
"""
# Grid is 1D, so each program instance processes one row.
pid_row = tl.program_id(0)
# Pointers to the current row for inputs and output.
row_hidden_states_ptr = hidden_states_ptr + pid_row * stride_hidden_states_batch
row_output_ptr = output_ptr + pid_row * stride_output_batch
# --- 1. Compute mean of squares ---
# Create a block of offsets for the hidden dimension.
offs_h = tl.arange(0, BLOCK_SIZE_H)
mask_h = offs_h < HIDDEN_SIZE
# Load the row of hidden_states. Use masking to handle cases where
# BLOCK_SIZE_H > HIDDEN_SIZE. `other=0.0` ensures that padding
# doesn't affect the sum of squares.
x_block = tl.load(row_hidden_states_ptr + offs_h, mask=mask_h, other=0.0)
# Promote to float32 for high-precision reduction.
x_f32 = x_block.to(tl.float32)
# Calculate sum of squares.
sum_sq = tl.sum(x_f32 * x_f32, axis=0)
# Calculate mean of squares. We divide by the actual HIDDEN_SIZE.
mean_sq = sum_sq / HIDDEN_SIZE
# --- 2. Compute inverse root mean square ---
# Add epsilon for numerical stability and compute rsqrt.
inv_rms = tl.rsqrt(mean_sq + EPS)
# --- 3. Normalize, scale, and store ---
# Load the corresponding weights.
weight_block = tl.load(weight_ptr + offs_h, mask=mask_h)
# Perform normalization and scaling.
# We reuse x_f32, which is already in registers from the initial load.
output_f32 = x_f32 * inv_rms * weight_block.to(tl.float32)
# Convert back to the output dtype (bfloat16) and store.
tl.store(row_output_ptr + offs_h, output_f32.to(tl.bfloat16), mask=mask_h)
def run(*args, **kwargs):
"""
Wrapper function for the RMSNorm Triton kernel.
Handles device management, argument parsing, and kernel launching.
It ensures that tensors are on the correct device (CUDA) for kernel
execution and that the output is moved back to the original device
of the input tensors.
Args:
hidden_states (torch.Tensor): The input tensor of shape [batch_size, 1536]
and dtype bfloat16.
weight (torch.Tensor): The weight tensor of shape [1536] and dtype bfloat16.
Returns:
torch.Tensor: The normalized and scaled output tensor with the same shape
and dtype as hidden_states.
"""
# --- Argument Parsing ---
hidden_states = kwargs.get('hidden_states')
weight = kwargs.get('weight')
if len(args) > 0:
if hidden_states is not None:
raise TypeError("run() got multiple values for argument 'hidden_states'")
hidden_states = args[0]
if len(args) > 1:
if weight is not None:
raise TypeError("run() got multiple values for argument 'weight'")
weight = args[1]
if len(args) > 2:
raise TypeError(f"run() takes 2 positional arguments but {len(args)} were given")
if hidden_states is None or weight is None:
raise TypeError("run() missing required arguments: 'hidden_states' or 'weight'")
# --- Input Validation ---
assert isinstance(hidden_states, torch.Tensor), "hidden_states must be a torch.Tensor"
assert isinstance(weight, torch.Tensor), "weight must be a torch.Tensor"
assert hidden_states.shape[1] == 1536, "hidden_size must be 1536"
assert hidden_states.shape[1] == weight.shape[0], "hidden_states and weight dimensions must match"
assert hidden_states.ndim == 2, "hidden_states must be a 2D tensor"
assert weight.ndim == 1, "weight must be a 1D tensor"
assert hidden_states.dtype == torch.bfloat16, "hidden_states dtype must be bfloat16"
assert weight.dtype == torch.bfloat16, "weight dtype must be bfloat16"
# --- Device Management ---
original_device = hidden_states.device
if not torch.cuda.is_available():
if original_device.type == 'cuda':
raise RuntimeError("CUDA is not available, but input tensors are on a CUDA device.")
# As Triton is unavailable, fall back to a native PyTorch implementation on CPU.
print("Warning: Triton requires a CUDA-enabled GPU. Falling back to reference implementation on CPU.")
return _reference_run(hidden_states, weight)
device = torch.device("cuda")
hidden_states_gpu = hidden_states.to(device)
weight_gpu = weight.to(device)
# --- Allocate Output Tensor ---
output_gpu = torch.empty_like(hidden_states_gpu)
# --- Kernel Launch ---
batch_size, hidden_size = hidden_states_gpu.shape
# Use a block size that is the next power of 2 to allow for efficient hardware utilization
BLOCK_SIZE_H = triton.next_power_of_2(hidden_size)
grid = (batch_size, )
_rmsnorm_h1536_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,
EPS=1e-6,
BLOCK_SIZE_H=BLOCK_SIZE_H,
)
# --- Move Result to Original Device ---
return output_gpu.to(original_device)scrolls · 180 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON