gemini-2.5-pro / tritondc28mj
gemini-2.5-pro_triton_dc28mj · gemini-2.5-pro · triton · Apache-2.0
Kernel source · 169 lines ↓holds 14 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 169 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-dc28mj?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16
Benchmark evidence
14 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 14 measurements ›Showing all 14 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:14fb30ecff672694407ade17e0e63783a7fe2acab06591fca11b38f132c9af44
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({'BLOCK_SIZE_H': 1024}, num_warps=4),Kernel source
main.py169 lines
import torch
import triton
import triton.language as tl
import math
@triton.autotune(
configs=[
triton.Config({'BLOCK_SIZE_H': 1024}, num_warps=4),
triton.Config({'BLOCK_SIZE_H': 2048}, num_warps=8),
triton.Config({'BLOCK_SIZE_H': 4096}, num_warps=16),
],
key=['HIDDEN_SIZE'],
)
@triton.jit
def _fused_add_rmsnorm_h4096_kernel(
# Pointers to tensors
hidden_states_ptr,
residual_ptr,
weight_ptr,
output_ptr,
# Other parameters
batch_size,
# Constants
HIDDEN_SIZE: tl.constexpr,
EPS: tl.constexpr,
BLOCK_SIZE_H: tl.constexpr,
):
"""
Triton kernel for fused add and RMSNorm.
Each program instance processes a single row of the input tensors.
"""
# Get the row index for the current program
row_idx = tl.program_id(0)
# --- Pointers for the current row ---
# Note: This assumes inputs are contiguous, which is a common case.
row_hidden_states_ptr = hidden_states_ptr + row_idx * HIDDEN_SIZE
row_residual_ptr = residual_ptr + row_idx * HIDDEN_SIZE
row_output_ptr = output_ptr + row_idx * HIDDEN_SIZE
# Weight pointer is the same for all rows
weight_base_ptr = weight_ptr
# --- Pass 1: Calculate the sum of squares for RMSNorm ---
# Accumulator for the variance, initialized to zero.
# We use float32 for precision.
var_accumulator = tl.zeros([1], dtype=tl.float32)
# Iterate over the hidden dimension in blocks of size BLOCK_SIZE_H
for h_offset in range(0, HIDDEN_SIZE, BLOCK_SIZE_H):
# Create a vector of offsets for the current block
h_offsets = h_offset + tl.arange(0, BLOCK_SIZE_H)
# Create a mask to handle the last block if HIDDEN_SIZE is not a multiple of BLOCK_SIZE_H
mask = h_offsets < HIDDEN_SIZE
# Load one block of hidden_states and residual
hidden_states_block = tl.load(row_hidden_states_ptr + h_offsets, mask=mask, other=0.0)
residual_block = tl.load(row_residual_ptr + h_offsets, mask=mask, other=0.0)
# Perform the addition: x = hidden_states + residual
# Cast to float32 for high-precision computation
x = hidden_states_block.to(tl.float32) + residual_block.to(tl.float32)
# Accumulate the sum of squares of x
var_accumulator += tl.sum(x * x, axis=0)
# After iterating through all blocks, compute the mean and inverse RMS
mean_var = var_accumulator / HIDDEN_SIZE
inv_rms = tl.rsqrt(mean_var + EPS)
# --- Pass 2: Apply normalization and scaling, and store the result ---
# Re-iterate over the hidden dimension to apply the calculated inv_rms
for h_offset in range(0, HIDDEN_SIZE, BLOCK_SIZE_H):
h_offsets = h_offset + tl.arange(0, BLOCK_SIZE_H)
mask = h_offsets < HIDDEN_SIZE
# Reload the input blocks for this pass
hidden_states_block = tl.load(row_hidden_states_ptr + h_offsets, mask=mask, other=0.0)
residual_block = tl.load(row_residual_ptr + h_offsets, mask=mask, other=0.0)
# Recompute x, same as in Pass 1
x = hidden_states_block.to(tl.float32) + residual_block.to(tl.float32)
# Load the corresponding block of weights
weight_block = tl.load(weight_base_ptr + h_offsets, mask=mask, other=0.0)
# Apply the normalization and scaling
normalized_x = x * inv_rms
output_block = normalized_x * weight_block.to(tl.float32)
# Convert the final result back to bfloat16 and store it in the output tensor
tl.store(row_output_ptr + h_offsets, output_block.to(tl.bfloat16), mask=mask)
def fused_add_rmsnorm_h4096(hidden_states: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""
Wrapper function for the fused_add_rmsnorm_h4096 Triton kernel.
Args:
hidden_states (torch.Tensor): Input tensor of shape [batch_size, 4096] and dtype bfloat16.
residual (torch.Tensor): Residual tensor of shape [batch_size, 4096] and dtype bfloat16.
weight (torch.Tensor): Weight tensor of shape [4096] and dtype bfloat16.
Returns:
torch.Tensor: The output tensor of shape [batch_size, 4096] and dtype bfloat16.
"""
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
original_device = hidden_states.device
is_cpu = original_device.type == 'cpu'
# Move tensors to GPU if they are on CPU
if is_cpu:
hidden_states = hidden_states.cuda()
residual = residual.cuda()
weight = weight.cuda()
# --- Input Validation ---
batch_size, hidden_size = hidden_states.shape
if hidden_size != 4096:
raise ValueError(f"hidden_size must be 4096, but got {hidden_size}")
if hidden_states.shape != residual.shape:
raise ValueError("hidden_states and residual must have the same shape.")
if weight.shape != (hidden_size,):
raise ValueError(f"weight must have shape [{hidden_size}], but got {weight.shape}")
expected_dtype = torch.bfloat16
if hidden_states.dtype != expected_dtype or residual.dtype != expected_dtype or weight.dtype != expected_dtype:
raise TypeError(f"All input tensors must have dtype {expected_dtype}.")
# --- Kernel Launch ---
# Allocate output tensor on the same device as the inputs
output = torch.empty_like(hidden_states)
# The grid is 1D, with one program instance per row in the batch
grid = (batch_size,)
_fused_add_rmsnorm_h4096_kernel[grid](
hidden_states,
residual,
weight,
output,
batch_size,
HIDDEN_SIZE=hidden_size,
EPS=1e-5,
)
# --- Finalization ---
# Move output back to the original device if necessary
if is_cpu:
output = output.cpu()
return output
def run(*args, **kwargs):
"""
Public entry point for the kernel. Handles flexible argument passing.
This function allows calling with either positional or keyword arguments.
"""
# A robust way to handle both *args and **kwargs by forwarding them
# to the main function.
if args:
return fused_add_rmsnorm_h4096(*args, **kwargs)
else:
return fused_add_rmsnorm_h4096(**kwargs)
scrolls · 169 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON