gemini-2.5-pro / tritonnicff0
gemini-2.5-pro_triton_nicff0 · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 143 lines ↓holds 3 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 143 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-nicff0?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:9e6abc101857a5b48b54bc94b2ddfae7697626d445ca6570f5819fe98f96489e
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 = 4
num_warps=4 # A good default for a block size of 512.Kernel source
main.py143 lines
import torch
import triton
import triton.language as tl
import math
# No hexadecimal float literals are used in this file.
@triton.jit
def rmsnorm_h512_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
stride_hidden_states_batch,
stride_output_batch,
HIDDEN_SIZE: tl.constexpr,
EPS: tl.constexpr,
):
"""
Triton kernel for RMSNorm.
Each program instance processes a single row of the input tensor.
"""
# The grid is 1D, so we can use `tl.program_id(0)` to get the row index.
pid = tl.program_id(axis=0)
# Pointers to the current row for inputs and outputs.
row_hidden_states_ptr = hidden_states_ptr + pid * stride_hidden_states_batch
row_output_ptr = output_ptr + pid * stride_output_batch
# Load the entire row of hidden_states. Since HIDDEN_SIZE is 512,
# a single program instance can handle it efficiently.
offsets = tl.arange(0, HIDDEN_SIZE)
# Load input data and promote to float32 for high-precision computation.
# This is crucial for the stability of the variance calculation.
x = tl.load(row_hidden_states_ptr + offsets, mask=offsets < HIDDEN_SIZE).to(tl.float32)
# --- Start of RMSNorm computation ---
# 1. Calculate the sum of squares.
# This is a block-level reduction performed efficiently by `tl.sum`.
sum_of_squares = tl.sum(x * x, axis=0)
# 2. Compute the mean and add epsilon.
mean_of_squares = sum_of_squares / HIDDEN_SIZE
# 3. Calculate the reciprocal square root (rsqrt) for normalization.
inv_rms = tl.rsqrt(mean_of_squares + EPS)
# 4. Normalize the input row.
# The `inv_rms` is a scalar, which is automatically broadcasted across the vector `x`.
normalized_x = x * inv_rms
# 5. Scale by the learnable weight parameter.
# Load the weights and promote to float32 for the multiplication.
w = tl.load(weight_ptr + offsets, mask=offsets < HIDDEN_SIZE).to(tl.float32)
output_f32 = normalized_x * w
# --- End of RMSNorm computation ---
# Cast back to the original bfloat16 dtype and store the result.
output = output_f32.to(tl.bfloat16)
tl.store(row_output_ptr + offsets, output, mask=offsets < HIDDEN_SIZE)
def run(hidden_states: torch.Tensor, weight: torch.Tensor, *args, **kwargs) -> torch.Tensor:
"""
Wrapper function for the RMSNorm Triton kernel.
This function handles device management, input validation, kernel launching,
and ensures the output tensor is on the same device as the input.
Args:
hidden_states (torch.Tensor): Input tensor of shape [batch_size, 512] and dtype bfloat16.
weight (torch.Tensor): Weight tensor of shape [512] and dtype bfloat16.
*args: Additional positional arguments (ignored).
**kwargs: Additional keyword arguments (ignored).
Returns:
torch.Tensor: The output tensor of the same shape and dtype as `hidden_states`.
"""
# --- Device Management & Input Validation ---
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
# Validate tensor shapes and dtypes as per the specification.
assert hidden_states.shape[1] == 512, f"Expected hidden_size=512, but got {hidden_states.shape[1]}"
assert weight.shape == (512,), f"Expected weight shape (512,), but got {weight.shape}"
assert hidden_states.dtype == torch.bfloat16, f"Expected hidden_states dtype bfloat16, but got {hidden_states.dtype}"
assert weight.dtype == torch.bfloat16, f"Expected weight dtype bfloat16, but got {weight.dtype}"
# Preserve the original device of the input tensor.
original_device = hidden_states.device
is_cpu_input = original_device.type == 'cpu'
# Move tensors to GPU if they are not already there.
if is_cpu_input:
hidden_states_gpu = hidden_states.cuda()
weight_gpu = weight.cuda()
elif original_device.type != 'cuda':
# Handle cases like 'mps' or other devices by moving data to cuda for computation
hidden_states_gpu = hidden_states.to('cuda')
weight_gpu = weight.to('cuda')
else:
hidden_states_gpu = hidden_states
weight_gpu = weight
# --- Kernel Launch ---
batch_size, hidden_size = hidden_states_gpu.shape
# Allocate the output tensor on the GPU.
output_gpu = torch.empty_like(hidden_states_gpu)
# Ensure tensors are contiguous for optimal memory access patterns.
# Triton kernels often assume dense, C-style memory layout.
hidden_states_gpu = hidden_states_gpu.contiguous()
weight_gpu = weight_gpu.contiguous()
output_gpu = output_gpu.contiguous()
# The grid is 1D, with one program instance per row in the batch.
grid = (batch_size,)
# The epsilon value for numerical stability.
EPS = 1e-6
# Launch the Triton kernel.
rmsnorm_h512_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, # Pass as a compile-time constant for optimization.
EPS=EPS, # Pass as a compile-time constant.
num_warps=4 # A good default for a block size of 512.
)
# --- Result Handling ---
# If the original input was on the CPU, move the result back to the CPU.
if is_cpu_input:
return output_gpu.to(original_device)
return output_gpu
scrolls · 143 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON