Skip to content
KernelIndex
Search⌘K

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-epilogues_with_bias = s + bias
mmaacc_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