gemini-2.5-pro / tritonxtl8hx
gemini-2.5-pro_triton_xtl8hx · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 179 lines ↓holds 3 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 179 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-xtl8hx?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:7ce758c012140b71baa84335500fbc0667e43d7add4fa05e47dbffaea84ef696
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.py179 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def fused_add_rmsnorm_h2048_kernel(
# Pointers to tensors
hidden_states_ptr,
residual_ptr,
weight_ptr,
output_ptr,
# Stride variables for memory access
stride_hidden_states_batch,
stride_residual_batch,
stride_output_batch,
# Other parameters
hidden_size,
# Constants
EPS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""
Triton kernel for fused Add + RMSNorm optimized for B200.
- Each program instance processes one row of the input tensors.
- It fuses the element-wise addition with the RMS normalization.
- Computation is done in float32 for precision, while I/O is bfloat16.
"""
# -----------------------------------------------------------
# Map program ids `pid` to the batch dimension.
# -----------------------------------------------------------
# Each program instance handles one row.
row_idx = tl.program_id(axis=0)
# -----------------------------------------------------------
# Pointers to the current row
# -----------------------------------------------------------
hidden_states_row_ptr = hidden_states_ptr + row_idx * stride_hidden_states_batch
residual_row_ptr = residual_ptr + row_idx * stride_residual_batch
output_row_ptr = output_ptr + row_idx * stride_output_batch
# -----------------------------------------------------------
# Load the row of data and compute the sum of squares
# -----------------------------------------------------------
# Create a block of offsets for the hidden dimension.
# Since BLOCK_SIZE is expected to be equal to hidden_size, we load the whole row.
offs = tl.arange(0, BLOCK_SIZE)
# Load the input row vectors, casting to float32 for computation.
hidden_states = tl.load(hidden_states_row_ptr + offs).to(tl.float32)
residual = tl.load(residual_row_ptr + offs).to(tl.float32)
# Fused add operation
x = hidden_states + residual
# Compute sum of squares for the variance calculation.
# This is a scalar value for the row after the reduction.
var = tl.sum(x * x, axis=0)
# -----------------------------------------------------------
# Compute RMS and apply normalization
# -----------------------------------------------------------
# Calculate inverse root mean square.
rstd = tl.rsqrt(var / hidden_size + EPS)
# Load the weight vector. It is broadcasted across all rows.
weight = tl.load(weight_ptr + offs).to(tl.float32)
# Normalize x and apply the learned scaling (weight).
output_f32 = x * rstd * weight
# -----------------------------------------------------------
# Write the output
# -----------------------------------------------------------
# Cast back to the output dtype (bfloat16) and store.
tl.store(output_row_ptr + offs, output_f32.to(tl.bfloat16))
def run(*args, **kwargs):
"""
Wrapper function for the fused_add_rmsnorm_h2048 Triton kernel.
Handles device management, tensor validation, grid computation, and kernel launch.
It moves tensors to the GPU, runs the kernel, and returns the result on the
original device of the first input tensor.
Args:
hidden_states (torch.Tensor): The main input tensor of shape [batch_size, 2048] and dtype bfloat16.
residual (torch.Tensor): The residual tensor to be added, with the same shape and dtype as hidden_states.
weight (torch.Tensor): The scaling weights of shape [2048] and dtype bfloat16.
Returns:
torch.Tensor: The output tensor with the same shape and dtype as hidden_states.
"""
# 1. Resolve and validate arguments from args and kwargs
# This allows for flexible calling conventions (positional or keyword).
arg_names = ['hidden_states', 'residual', 'weight']
if args:
if len(args) > len(arg_names):
raise TypeError(f"run() takes at most {len(arg_names)} positional arguments but {len(args)} were given")
for i, arg in enumerate(args):
kwargs[arg_names[i]] = arg
hidden_states = kwargs.get('hidden_states')
residual = kwargs.get('residual')
weight = kwargs.get('weight')
if hidden_states is None or residual is None or weight is None:
missing = [name for name in arg_names if name not in kwargs]
raise TypeError(f"run() missing required arguments: {', '.join(missing)}")
# 2. Device Management: determine target device and move tensors
if not torch.cuda.is_available():
raise RuntimeError("Triton kernel requires a CUDA-enabled GPU.")
initial_devices = {
'hidden_states': hidden_states.device,
'residual': residual.device,
'weight': weight.device
}
# Determine the target CUDA device. If any tensor is on CUDA, use that device.
# Otherwise, default to the current CUDA device.
target_device = None
for tensor in [hidden_states, residual, weight]:
if tensor.is_cuda:
if target_device is None:
target_device = tensor.device
elif target_device != tensor.device:
raise ValueError("All input tensors must be on the same CUDA device.")
if target_device is None:
target_device = torch.device("cuda")
# Move all tensors to the target device for the kernel execution.
hidden_states_gpu = hidden_states.to(target_device)
residual_gpu = residual.to(target_device)
weight_gpu = weight.to(target_device)
# 3. Shape and DType validation on the device
B, H = hidden_states_gpu.shape
assert H == 2048, f"Expected hidden_size=2048, but got {H}"
assert hidden_states_gpu.shape == residual_gpu.shape, "hidden_states and residual must have the same shape"
assert weight_gpu.shape == (H,), f"Expected weight shape ({H},), but got {weight_gpu.shape}"
assert hidden_states_gpu.ndim == 2, "Inputs must be 2D tensors"
for name, tensor in [('hidden_states', hidden_states_gpu), ('residual', residual_gpu), ('weight', weight_gpu)]:
if tensor.dtype != torch.bfloat16:
raise TypeError(f"Expected {name} to be bfloat16, but got {tensor.dtype}")
# 4. Allocate output tensor on the target device
output = torch.empty_like(hidden_states_gpu)
# 5. Set up grid and launch the kernel
grid = (B,)
# B200 benefits from more warps for large memory-bound operations.
# For a large block size of 2048, 8 warps is a good choice.
num_warps = 8
fused_add_rmsnorm_h2048_kernel[grid](
hidden_states_gpu,
residual_gpu,
weight_gpu,
output,
hidden_states_gpu.stride(0),
residual_gpu.stride(0),
output.stride(0),
H,
EPS=1e-6,
BLOCK_SIZE=H,
num_warps=num_warps
)
# 6. Move the result back to the original device of the primary input tensor
output_device = initial_devices['hidden_states']
return output.to(output_device)scrolls · 179 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON