gemini-2.5-pro / triton0b5fbf
gemini-2.5-pro_triton_0b5fbf · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 254 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-0b5fbf?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, fp8_e4m3, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:b55a291b577a8a55bdaa5a8e82a949c6526c9a62efa0b5e72ea83d11cb885eb2
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.
fused-epilogue
s_with_bias = s + biasmma
acc_x1 += tl.dot(w1_x1_fp8.to(tl.float32) * w1_s1, a_tile)Kernel source
main.py254 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def moe_fp8_block_scale_ds_routing_topk8_ng8_kg4_e32_h7168_i2048_kernel(
# Pointers to Tensors
routing_logits_ptr,
routing_bias_ptr,
hidden_states_ptr,
hidden_states_scale_ptr,
gemm1_weights_ptr,
gemm1_weights_scale_ptr,
gemm2_weights_ptr,
gemm2_weights_scale_ptr,
output_ptr,
# Scalar Arguments
local_expert_offset,
routed_scaling_factor,
seq_len,
# Strides
stride_logits_s, stride_logits_e,
stride_bias_e,
stride_hidden_s, stride_hidden_h,
stride_h_scale_h, stride_h_scale_s,
stride_w1_e, stride_w1_g, stride_w1_h,
stride_w1_scale_e, stride_w1_scale_g, stride_w1_scale_h,
stride_w2_e, stride_w2_h, stride_w2_i,
stride_w2_scale_e, stride_w2_scale_h, stride_w2_scale_i,
stride_out_s, stride_out_h,
# Compile-time Constants
NUM_EXPERTS: tl.constexpr,
NUM_LOCAL_EXPERTS: tl.constexpr,
HIDDEN_SIZE: tl.constexpr,
INTERMEDIATE_SIZE: tl.constexpr,
GEMM1_OUT_SIZE: tl.constexpr,
TOP_K: tl.constexpr,
N_GROUP: tl.constexpr,
TOPK_GROUP: tl.constexpr,
GROUP_SIZE: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
# Tiling configuration for GEMMs
BLOCK_K_GEMM1: tl.constexpr,
BLOCK_I: tl.constexpr,
BLOCK_H: tl.constexpr,
):
# Each program instance computes one token.
pid = tl.program_id(0)
# Constants
NEG_INF = -float('inf')
# --- 1. On-Chip Routing Logic ---
e_range = tl.arange(0, NUM_EXPERTS)
# Load logits and bias for the current token
logits_ptr = routing_logits_ptr + pid * stride_logits_s
logits = tl.load(logits_ptr + e_range * stride_logits_e).to(tl.float32)
s = tl.sigmoid(logits)
bias = tl.load(routing_bias_ptr + e_range * stride_bias_e).to(tl.bfloat16).to(tl.float32)
s_with_bias = s + bias
# [FIXED] Group scores: sum of top-2 in each group
s_wb_grouped = tl.reshape(s_with_bias, (N_GROUP, GROUP_SIZE))
top1_per_group = tl.reduce(s_wb_grouped, axis=1, combine_fn=tl.max)
s_wb_masked_1 = tl.where(s_wb_grouped == top1_per_group[:, None], NEG_INF, s_wb_grouped)
top2_per_group = tl.reduce(s_wb_masked_1, axis=1, combine_fn=tl.max)
group_scores = top1_per_group + top2_per_group
# Select top-k groups using iterative find-max-and-mask
selected_group_scores = group_scores
top_group_indices = tl.zeros((TOPK_GROUP,), dtype=tl.int32)
for i in tl.static_range(TOPK_GROUP):
idx = tl.argmax(selected_group_scores, axis=0)
top_group_indices = tl.where(tl.arange(0, TOPK_GROUP) == i, idx, top_group_indices)
selected_group_scores = tl.where(tl.arange(0, N_GROUP) == idx, NEG_INF, selected_group_scores)
# Create mask for experts in selected groups
group_mask = tl.zeros((NUM_EXPERTS,), dtype=tl.int1)
for i in tl.static_range(TOPK_GROUP):
g_idx = top_group_indices[i]
start, end = g_idx * GROUP_SIZE, (g_idx + 1) * GROUP_SIZE
group_mask = tl.where((e_range >= start) & (e_range < end), 1, group_mask)
scores_pruned = tl.where(group_mask, s_with_bias, NEG_INF)
# Global top-k experts from pruned scores
selected_expert_scores = scores_pruned
topk_indices = tl.zeros((TOP_K,), dtype=tl.int32)
for i in tl.static_range(TOP_K):
idx = tl.argmax(selected_expert_scores, axis=0)
topk_indices = tl.where(tl.arange(0, TOP_K) == i, idx, topk_indices)
selected_expert_scores = tl.where(e_range == idx, NEG_INF, selected_expert_scores)
# Calculate final routing weights
weights_mask = tl.zeros((NUM_EXPERTS,), dtype=tl.int1)
for i in tl.static_range(TOP_K):
weights_mask = tl.where(e_range == topk_indices[i], 1, weights_mask)
weights = tl.where(weights_mask, s, 0.0)
weights_sum = tl.sum(weights, axis=0)
weights = weights / (weights_sum + 1e-20) * routed_scaling_factor
# --- 2. Dequantize Input Hidden States (A) ---
h_offsets_full = tl.arange(0, HIDDEN_SIZE)
h_block_indices = h_offsets_full // BLOCK_SIZE
a_fp8 = tl.load(hidden_states_ptr + pid * stride_hidden_s + h_offsets_full * stride_hidden_h)
a_scales = tl.load(hidden_states_scale_ptr + h_block_indices * stride_h_scale_h + pid * stride_h_scale_s)
a_dequant = a_fp8.to(tl.float32) * a_scales
# --- 3. Expert Computation and Accumulation (Tiled over H) ---
for h_base in range(0, HIDDEN_SIZE, BLOCK_H):
h_offsets = h_base + tl.arange(0, BLOCK_H)
h_mask = h_offsets < HIDDEN_SIZE
final_output_tile = tl.zeros((BLOCK_H,), dtype=tl.float32)
for k_expert_idx in tl.static_range(TOP_K):
ge = topk_indices[k_expert_idx]
is_local = (ge >= local_expert_offset) & (ge < local_expert_offset + NUM_LOCAL_EXPERTS)
if is_local:
le = ge - local_expert_offset
weight = weights[ge]
expert_output_tile = tl.zeros((BLOCK_H,), dtype=tl.float32)
# Loop over intermediate size (K dimension of GEMM2)
for i_base in range(0, INTERMEDIATE_SIZE, BLOCK_I):
i_offsets = i_base + tl.arange(0, BLOCK_I)
i_mask = i_offsets < INTERMEDIATE_SIZE
# --- Step 1: Compute C_tile = SwiGLU(A @ W13_tile.T) ---
acc_x1 = tl.zeros((BLOCK_I,), dtype=tl.float32)
acc_x2 = tl.zeros((BLOCK_I,), dtype=tl.float32)
for k1_base in range(0, HIDDEN_SIZE, BLOCK_K_GEMM1):
k1_offsets = k1_base + tl.arange(0, BLOCK_K_GEMM1)
k1_mask = k1_offsets < HIDDEN_SIZE
a_tile = tl.load(a_dequant + k1_offsets, mask=k1_mask, other=0.0)
k1_block_idx = k1_base // BLOCK_SIZE
# Process X1 (gate)
w1_x1_ptr = gemm1_weights_ptr + le*stride_w1_e + i_offsets[:,None]*stride_w1_g + k1_offsets[None,:]*stride_w1_h
w1_x1_fp8 = tl.load(w1_x1_ptr, mask=i_mask[:, None] & k1_mask[None, :], other=0.0)
w1_s1_ptr = gemm1_weights_scale_ptr + le*stride_w1_scale_e + (i_base//BLOCK_SIZE)*stride_w1_scale_g + k1_block_idx*stride_w1_scale_h
w1_s1 = tl.load(w1_s1_ptr)
acc_x1 += tl.dot(w1_x1_fp8.to(tl.float32) * w1_s1, a_tile)
# Process X2 (up)
w1_x2_ptr = gemm1_weights_ptr + le*stride_w1_e + (i_offsets[:,None]+INTERMEDIATE_SIZE)*stride_w1_g + k1_offsets[None,:]*stride_w1_h
w1_x2_fp8 = tl.load(w1_x2_ptr, mask=i_mask[:, None] & k1_mask[None, :], other=0.0)
w1_s2_ptr = gemm1_weights_scale_ptr + le*stride_w1_scale_e + ((i_base+INTERMEDIATE_SIZE)//BLOCK_SIZE)*stride_w1_scale_g + k1_block_idx*stride_w1_scale_h
w1_s2 = tl.load(w1_s2_ptr)
acc_x2 += tl.dot(w1_x2_fp8.to(tl.float32) * w1_s2, a_tile)
# [FIXED] SwiGLU: C = X1 * silu(X2) = X1 * (X2 * sigmoid(X2))
silu_x2 = acc_x2 * tl.sigmoid(acc_x2)
c_tile = acc_x1 * silu_x2
# --- Step 2: Accumulate expert_output_tile += C_tile @ W2_tile.T ---
w2_ptr = gemm2_weights_ptr + le*stride_w2_e + h_offsets[:,None]*stride_w2_h + i_offsets[None,:]*stride_w2_i
w2_fp8 = tl.load(w2_ptr, mask=h_mask[:, None] & i_mask[None, :], other=0.0)
w2_s_ptr = gemm2_weights_scale_ptr + le*stride_w2_scale_e + (h_base//BLOCK_SIZE)*stride_w2_scale_h + (i_base//BLOCK_SIZE)*stride_w2_scale_i
w2_s = tl.load(w2_s_ptr)
w2_dequant = w2_fp8.to(tl.float32) * w2_s
expert_output_tile += tl.dot(w2_dequant, c_tile)
final_output_tile += expert_output_tile * weight
# --- 4. Store final result tile ---
out_ptr = output_ptr + pid * stride_out_s + h_offsets * stride_out_h
tl.store(out_ptr, final_output_tile, mask=h_mask)
def run(
routing_logits: torch.Tensor,
routing_bias: torch.Tensor,
hidden_states: torch.Tensor,
hidden_states_scale: torch.Tensor,
gemm1_weights: torch.Tensor,
gemm1_weights_scale: torch.Tensor,
gemm2_weights: torch.Tensor,
gemm2_weights_scale: torch.Tensor,
local_expert_offset: int,
routed_scaling_factor: float,
):
"""
Wrapper function to run the MoE Triton kernel with automatic device management.
"""
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
device = routing_logits.device
if device.type != 'cuda':
raise RuntimeError(f"Input tensors must be on a CUDA device, but found {device.type}.")
seq_len, num_experts = routing_logits.shape
hidden_size = hidden_states.shape[1]
output = torch.empty((seq_len, hidden_size), device=device, dtype=torch.float32)
grid = (seq_len,)
constants = {
"NUM_EXPERTS": 256,
"NUM_LOCAL_EXPERTS": 32,
"HIDDEN_SIZE": 7168,
"INTERMEDIATE_SIZE": 2048,
"GEMM1_OUT_SIZE": 4096,
"TOP_K": 8,
"N_GROUP": 8,
"TOPK_GROUP": 4,
"GROUP_SIZE": 256 // 8,
"BLOCK_SIZE": 128,
"BLOCK_K_GEMM1": 128,
"BLOCK_I": 64,
"BLOCK_H": 128,
}
inputs_to_check = [
routing_logits, routing_bias, hidden_states, hidden_states_scale,
gemm1_weights, gemm1_weights_scale, gemm2_weights, gemm2_weights_scale
]
contiguous_inputs = []
for t in inputs_to_check:
if t.device != device:
t = t.to(device)
if not t.is_contiguous():
t = t.contiguous()
contiguous_inputs.append(t)
(routing_logits, routing_bias, hidden_states, hidden_states_scale,
gemm1_weights, gemm1_weights_scale, gemm2_weights, gemm2_weights_scale) = contiguous_inputs
moe_fp8_block_scale_ds_routing_topk8_ng8_kg4_e32_h7168_i2048_kernel[grid](
routing_logits, routing_bias, hidden_states, hidden_states_scale,
gemm1_weights, gemm1_weights_scale, gemm2_weights, gemm2_weights_scale,
output,
local_expert_offset, routed_scaling_factor,
seq_len,
routing_logits.stride(0), routing_logits.stride(1),
routing_bias.stride(0),
hidden_states.stride(0), hidden_states.stride(1),
hidden_states_scale.stride(0), hidden_states_scale.stride(1),
gemm1_weights.stride(0), gemm1_weights.stride(1), gemm1_weights.stride(2),
gemm1_weights_scale.stride(0), gemm1_weights_scale.stride(1), gemm1_weights_scale.stride(2),
gemm2_weights.stride(0), gemm2_weights.stride(1), gemm2_weights.stride(2),
gemm2_weights_scale.stride(0), gemm2_weights_scale.stride(1), gemm2_weights_scale.stride(2),
output.stride(0), output.stride(1),
**constants
)
return output.to(torch.bfloat16)
scrolls · 254 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON