gemini-2.5-pro / triton5h17k3
gemini-2.5-pro_triton_5h17k3 · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 174 lines ↓holds 3 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 174 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-5h17k3?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:f6b8f61700fc285ee9b0aa0edcb2dbe9b3b78917612e7e0eec92f02b2a4f1159
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,Kernel source
main.py174 lines
import torch
import triton
import triton.language as tl
import math
# Reference implementation for mathematical specification verification
# @torch.no_grad()
# def reference_run(hidden_states, weight):
# batch_size, hidden_size = hidden_states.shape
# # Check constants
# assert hidden_size == 7168
#
# 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_kernel(
# Pointers to tensors
x_ptr,
weight_ptr,
output_ptr,
# Stride to move to the next row
stride_x_batch,
stride_out_batch,
# Matrix dimensions
hidden_size,
# Kernel constants
EPS: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
):
"""
Triton kernel for RMS Normalization.
This kernel is optimized for a fixed hidden_size and targets B200 GPUs.
Grid: 1D, with each program instance processing one row (one item in the batch).
"""
# Each program instance processes a single row.
pid_batch = tl.program_id(axis=0)
# Pointers to the current row for input and output.
row_x_ptr = x_ptr + pid_batch * stride_x_batch
row_output_ptr = output_ptr + pid_batch * stride_out_batch
# B200 Optimization: Use a large block size to process the entire row in a single,
# vectorized operation. This maximizes memory bandwidth utilization.
# `BLOCK_SIZE_N` is configured to be the next power of 2 of `hidden_size`.
offsets_n = tl.arange(0, BLOCK_SIZE_N)
mask_n = offsets_n < hidden_size
# --- Pass 1: Compute sum of squares and inv_rms ---
# Load the entire row into registers (SRAM).
# Convert to float32 for high-precision accumulation to avoid overflow/underflow.
x = tl.load(row_x_ptr + offsets_n, mask=mask_n, other=0.0).to(tl.float32)
# Compute the sum of squares. tl.sum performs an efficient reduction.
sum_sq = tl.sum(x * x, axis=0)
# Calculate variance and inverse root mean square.
var = sum_sq / hidden_size
inv_rms = tl.rsqrt(var + EPS)
# --- Pass 2: Normalize, scale, and store ---
# This pass is fused and operates on data held in registers.
# Load the corresponding weights.
w = tl.load(weight_ptr + offsets_n, mask=mask_n, other=0.0).to(tl.float32)
# Apply the normalization and scaling.
output_val = x * inv_rms * w
# Convert back to the target dtype (bfloat16) and store the result.
tl.store(row_output_ptr + offsets_n, output_val.to(tl.bfloat16), mask=mask_n)
def rmsnorm_h7168(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""
Wrapper for the RMSNorm Triton kernel.
Args:
hidden_states (torch.Tensor): Input tensor of shape [batch_size, 7168] and dtype bfloat16.
weight (torch.Tensor): Weight tensor of shape [7168] and dtype bfloat16.
Returns:
torch.Tensor: The normalized and scaled output tensor.
"""
# Input validation
if hidden_states.shape[1] != 7168:
raise ValueError(f"Expected hidden_size=7168, but got {hidden_states.shape[1]}")
if hidden_states.dtype != torch.bfloat16:
raise TypeError(f"Expected hidden_states dtype bfloat16, but got {hidden_states.dtype}")
if weight.shape != (7168,):
raise ValueError(f"Expected weight shape (7168,), but got {weight.shape}")
if weight.dtype != torch.bfloat16:
raise TypeError(f"Expected weight dtype bfloat16, but got {weight.dtype}")
# Kernel parameters
batch_size, hidden_size = hidden_states.shape
# Allocate output tensor
output = torch.empty_like(hidden_states)
# Grid definition: one program per row
grid = (batch_size,)
# B200 Optimization: Choose a block size that covers the entire row dimension.
# This allows for full vectorization and avoids loop overhead within the kernel.
BLOCK_SIZE_N = triton.next_power_of_2(hidden_size)
# Kernel launch
_rmsnorm_kernel[grid](
x_ptr=hidden_states,
weight_ptr=weight,
output_ptr=output,
stride_x_batch=hidden_states.stride(0),
stride_out_batch=output.stride(0),
hidden_size=hidden_size,
EPS=1e-6,
BLOCK_SIZE_N=BLOCK_SIZE_N,
# B200 Optimization: Use a higher number of warps to hide memory latency,
# which is crucial for memory-bound operations like this.
num_warps=8,
)
return output
def run(*args, **kwargs):
"""
Public entry point for the RMSNorm operation.
Handles device management and calls the Triton kernel implementation.
"""
# 1. Parse arguments
if args:
hidden_states, weight = args
elif kwargs:
hidden_states = kwargs.get('hidden_states')
weight = kwargs.get('weight')
else:
raise ValueError("Missing required arguments 'hidden_states' and 'weight'")
if hidden_states is None or weight is None:
raise ValueError("Both 'hidden_states' and 'weight' must be provided")
# 2. Device Management: Setup
original_device = hidden_states.device
if not torch.cuda.is_available():
raise RuntimeError("Triton requires CUDA, but torch.cuda.is_available() is False.")
target_device = torch.device("cuda")
# 3. Move tensors to GPU if they aren't already
inputs_on_gpu = True
if hidden_states.device != target_device:
hidden_states = hidden_states.to(target_device)
inputs_on_gpu = False
if weight.device != target_device:
weight = weight.to(target_device)
inputs_on_gpu = False
# 4. Execute the kernel
output = rmsnorm_h7168(hidden_states, weight)
# 5. Device Management: Move result back to the original device
if original_device != target_device:
output = output.to(original_device)
return output
scrolls · 174 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON