gemini-2.5-pro_triton_jaedrq
gemini-2.5-pro · triton · Apache-2.0
Kernel source · 163 lines ↓holds 2 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 163 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-jaedrq?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:8110ee2a517d85eeab6a1ac5b612cca680305217c75945cae1c18c770ce3e834
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,Kernel source
main.py163 lines
import torch
import triton
import triton.language as tl
import math
#
# Reference Implementation:
# import torch
#
# @torch.no_grad()
# def run(hidden_states, weight):
# batch_size, hidden_size = hidden_states.shape
# # Check constants
# assert hidden_size == 128
#
# 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_h128_kernel(
hidden_states_ptr,
weight_ptr,
output_ptr,
batch_size,
hidden_size,
stride_b,
EPS: tl.constexpr,
BLOCK_SIZE_H: tl.constexpr,
):
"""
Triton kernel for RMS Normalization with a fixed hidden_size of 128.
Each program in the grid processes one row from the batch.
"""
# Get the row index for the current program
pid_b = tl.program_id(axis=0)
# Create pointers to the start of the current row for inputs and output
row_x_ptr = hidden_states_ptr + pid_b * stride_b
row_y_ptr = output_ptr + pid_b * stride_b
# Create a range of offsets for the hidden dimension
# Since BLOCK_SIZE_H is fixed to hidden_size (128), we load the entire row
offsets_h = tl.arange(0, BLOCK_SIZE_H)
# Load the full row of hidden_states and the full weight vector
# No mask is needed as hidden_size == BLOCK_SIZE_H
x = tl.load(row_x_ptr + offsets_h)
w = tl.load(weight_ptr + offsets_h)
# --- Computation is performed in float32 for precision ---
x_fp32 = x.to(tl.float32)
w_fp32 = w.to(tl.float32)
# 1. Square the elements
x_sq = x_fp32 * x_fp32
# 2. Compute the mean of the squares (reduction)
# tl.sum performs an efficient reduction over the block of 128 elements
var = tl.sum(x_sq, axis=0) / hidden_size
# 3. Compute the inverse root mean square
inv_rms = tl.rsqrt(var + EPS)
# 4. Normalize the hidden states and apply the learned scaling factor (weight)
y = x_fp32 * inv_rms * w_fp32
# --- Cast back to bfloat16 and store the result ---
y_bf16 = y.to(tl.bfloat16)
tl.store(row_y_ptr + offsets_h, y_bf16)
def run(*args, **kwargs):
"""
Wrapper function to run the RMSNorm Triton kernel.
Handles device management, tensor validation, and kernel launching.
It preserves the device of the input tensors for the output.
Args:
hidden_states (torch.Tensor): Input tensor of shape [batch_size, 128]
and dtype bfloat16.
weight (torch.Tensor): Weight tensor of shape [128] and dtype bfloat16.
Returns:
torch.Tensor: The normalized output tensor of the same shape and dtype
as hidden_states.
"""
# 1. Parse arguments
if args:
if len(args) != 2:
raise ValueError(f"Expected 2 positional arguments, but got {len(args)}")
hidden_states, weight = args
else:
hidden_states = kwargs.get('hidden_states')
weight = kwargs.get('weight')
if hidden_states is None or weight is None:
raise ValueError("Missing required keyword arguments: 'hidden_states' and/or 'weight'")
# 2. Validate tensor properties
if hidden_states.dim() != 2 or hidden_states.shape[1] != 128:
raise ValueError(f"Expected hidden_states to have shape [batch_size, 128], but got {hidden_states.shape}")
if weight.dim() != 1 or weight.shape[0] != 128:
raise ValueError(f"Expected weight to have shape [128], but got {weight.shape}")
if hidden_states.dtype != torch.bfloat16:
raise TypeError(f"Expected hidden_states to have dtype torch.bfloat16, but got {hidden_states.dtype}")
if weight.dtype != torch.bfloat16:
raise TypeError(f"Expected weight to have dtype torch.bfloat16, but got {weight.dtype}")
# 3. Device management
original_device = hidden_states.device
is_cpu_input = original_device.type == 'cpu'
# If inputs are on CPU, they must be moved to a GPU to run the Triton kernel.
if is_cpu_input:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
target_device = 'cuda'
hidden_states = hidden_states.to(target_device)
weight = weight.to(target_device)
# If inputs are already on a GPU, ensure they are on the same device.
else:
if hidden_states.device != weight.device:
raise ValueError(f"hidden_states and weight must be on the same device, "
f"but got {hidden_states.device} and {weight.device}")
target_device = hidden_states.device
# 4. Prepare for kernel launch
batch_size, hidden_size = hidden_states.shape
# Allocate the output tensor on the target GPU device
output = torch.empty_like(hidden_states)
# The grid is 1D, with one program per row in the batch.
grid = (batch_size,)
# Constants for the kernel
EPS = 1e-6
BLOCK_SIZE_H = 128
# 5. Launch the Triton kernel
# num_warps=4 is a robust choice for a block size of 128 on modern GPUs like B200.
rmsnorm_h128_kernel[grid](
hidden_states,
weight,
output,
batch_size,
hidden_size,
hidden_states.stride(0),
EPS=EPS,
BLOCK_SIZE_H=BLOCK_SIZE_H,
num_warps=4,
)
# 6. Move the result back to the original device if necessary
if is_cpu_input:
output = output.to(original_device)
return outputscrolls · 163 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON