Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_triton_544238

gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 246 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-544238?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32, 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:c6003e0ad696af19229b4aa17de085c4e898f73dd3cd017481fc579b918b9cf7
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.

tile-m = 512BLOCK_SIZE_M = 512

Kernel source

main.py246 lines
import torch
import triton
import triton.language as tl
import math

# B200 is part of the Blackwell architecture. Optimizations for Hopper
# (large SRAM, efficient block-level primitives) are expected to perform well on B200.
# This kernel is redesigned using modern Triton features to be correct and efficient.

# --- Triton Kernel ---

@triton.jit
def top_k_top_p_sampling_from_probs_v128256_kernel(
    probs_ptr,
    top_k_ptr,
    top_p_ptr,
    samples_ptr,
    seed_ptr,
    batch_size,
    stride_probs_b,
    VOCAB_SIZE: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    TOP_K_BUFFER_SIZE: tl.constexpr,
):
    """
    Triton kernel for Top-K, Top-P sampling.

    This kernel processes one sequence per program instance (one row of the batch).
    It uses a streaming reduction to find the top candidates from the vocabulary,
    applies top-k and top-p filtering, and finally performs multinomial sampling.
    This version uses modern Triton APIs and direct operations on tl.tensor objects.

    - VOCAB_SIZE: The total vocabulary size.
    - BLOCK_SIZE_M: The size of blocks to read from the vocab. Must be a power of 2.
    - TOP_K_BUFFER_SIZE: The size of the SRAM buffer for holding top candidates. Must be a power of 2.
    """
    pid = tl.program_id(0)
    if pid >= batch_size:
        return

    # --- Load per-sequence parameters ---
    k_val = tl.load(top_k_ptr + pid)
    p_val = tl.load(top_p_ptr + pid)
    seed = tl.load(seed_ptr + pid)

    probs_row_ptr = probs_ptr + pid * stride_probs_b

    # The size of our merge buffer for the streaming top-k reduction.
    # Must be a power of 2 for tl.sort.
    # FIX: Declared as tl.constexpr to resolve the compilation error with tl.arange.
    MERGE_BUFFER_SIZE: tl.constexpr = TOP_K_BUFFER_SIZE + BLOCK_SIZE_M

    # --- SRAM Allocation for top candidates (as registers) ---
    sram_top_k_probs = tl.full([TOP_K_BUFFER_SIZE], -1.0, dtype=tl.float32)
    sram_top_k_indices = tl.full([TOP_K_BUFFER_SIZE], -1, dtype=tl.int32)

    # --- Streaming Top-K Reduction ---
    # Find the top `TOP_K_BUFFER_SIZE` candidates from the entire vocabulary.
    num_blocks = tl.cdiv(VOCAB_SIZE, BLOCK_SIZE_M)
    for block_idx in range(num_blocks):
        # Load a block of probabilities and their corresponding indices from HBM
        m_offsets = tl.arange(0, BLOCK_SIZE_M)
        current_offsets = block_idx * BLOCK_SIZE_M + m_offsets
        mask = current_offsets < VOCAB_SIZE

        chunk_probs = tl.load(probs_row_ptr + current_offsets, mask=mask, other=-1.0)
        chunk_indices = current_offsets

        # --- Merge and Sort in Registers/SRAM ---
        # 1. Construct the merged buffer of candidates by concatenating the
        #    current top-k with the new chunk.
        # FIX: The original indexing logic was flawed and caused out-of-bounds access.
        # This corrected version clamps indices to be safe for both branches of tl.where.
        merged_offsets = tl.arange(0, MERGE_BUFFER_SIZE)
        is_top_k_part = merged_offsets < TOP_K_BUFFER_SIZE

        # Safely clamp indices for the SRAM part to [0, TOP_K_BUFFER_SIZE - 1]
        sram_indices_safe = tl.minimum(merged_offsets, TOP_K_BUFFER_SIZE - 1)
        # Safely clamp indices for the chunk part to [0, BLOCK_SIZE_M - 1]
        chunk_indices_safe = tl.maximum(0, merged_offsets - TOP_K_BUFFER_SIZE)
        chunk_indices_safe = tl.minimum(chunk_indices_safe, BLOCK_SIZE_M - 1)

        merged_probs = tl.where(is_top_k_part, sram_top_k_probs[sram_indices_safe], chunk_probs[chunk_indices_safe])
        merged_indices = tl.where(is_top_k_part, sram_top_k_indices[sram_indices_safe], chunk_indices[chunk_indices_safe])

        # 2. Pack probs (key) and indices (value) into int64 for a single sort operation.
        #    To sort floats in descending order, we negate their integer representation.
        probs_as_int = merged_probs.to(tl.int32, bitcast=True)
        neg_probs_as_int = -probs_as_int
        packed_data = neg_probs_as_int.to(tl.int64) << 32 | merged_indices.to(tl.int64)

        # 3. Sort the packed data. tl.sort is a highly optimized block-level primitive.
        sorted_packed = tl.sort(packed_data)

        # 4. Unpack the data and update the top-k buffers for the next iteration.
        k_offsets = tl.arange(0, TOP_K_BUFFER_SIZE)
        top_k_packed_slice = sorted_packed[k_offsets]
        
        unpacked_neg_probs_as_int = (top_k_packed_slice >> 32).to(tl.int32)
        sram_top_k_indices = (top_k_packed_slice & 0xFFFFFFFF).to(tl.int32)
        sram_top_k_probs = (-unpacked_neg_probs_as_int).to(tl.float32, bitcast=True)

    # `sram_top_k_probs` and `sram_top_k_indices` now hold the top candidates.

    # --- Apply Top-K filtering ---
    num_candidates = TOP_K_BUFFER_SIZE
    if 0 < k_val < VOCAB_SIZE:
        num_candidates = tl.minimum(k_val, TOP_K_BUFFER_SIZE)

    cand_offsets = tl.arange(0, TOP_K_BUFFER_SIZE)
    k_mask = cand_offsets < num_candidates
    
    # --- Greedy sampling (p <= 0.0) ---
    if p_val <= 0.0:
        # The candidates are sorted, so the first element is the argmax.
        result_idx = sram_top_k_indices[0]
        tl.store(samples_ptr + pid, result_idx.to(tl.int64))
        return

    # --- Apply Top-P (Nucleus) filtering ---
    candidate_probs = tl.where(k_mask, sram_top_k_probs, 0.0)
    
    if p_val < 1.0:
        # Renormalize the candidate probabilities before calculating CDF for top-p.
        total_prob_sum = tl.sum(candidate_probs, axis=0)
        # Avoid division by zero if all candidate probs are zero
        if total_prob_sum > 1e-9:
            probs_for_p = candidate_probs / total_prob_sum
            
            # Compute CDF. Keep token `i` if `cdf[i-1] <= p`.
            # This is equivalent to `(cumsum(p) - p) <= p`.
            cdf = tl.cumsum(probs_for_p, axis=0)
            shifted_cdf = cdf - probs_for_p
            
            p_mask = shifted_cdf < p_val
            candidate_probs = tl.where(p_mask, candidate_probs, 0.0)
        else:
            # If sum is zero, all candidate_probs are already zero, so do nothing.
            pass

    # --- Final Multinomial Sampling ---
    final_probs = candidate_probs
    final_prob_sum = tl.sum(final_probs, axis=0)
    
    result_idx = -1
    if final_prob_sum > 1e-9: # Use a small epsilon for float comparison
        # Generate a random number in [0, 1) and scale it.
        rand_offset = pid # Use a unique offset for per-row randomness
        r_val = tl.rand(seed, rand_offset)
        sample_thresh = r_val * final_prob_sum
        
        sample_cdf = tl.cumsum(final_probs, axis=0)
        
        # Find the first index `i` where `sample_cdf[i] > sample_thresh`.
        is_winner = sample_cdf > sample_thresh
        # tl.argmax returns the index of the first '1'
        winner_sram_idx = tl.argmax(is_winner.to(tl.int32), axis=0)
        result_idx = sram_top_k_indices[winner_sram_idx]
    else:
        # If all probabilities are filtered out (e.g., k=0 or p is very small),
        # fall back to the absolute top token (greedy).
        result_idx = sram_top_k_indices[0]

    tl.store(samples_ptr + pid, result_idx.to(tl.int64))


def run(*args, **kwargs):
    """
    Wrapper function for the Top-K, Top-P sampling Triton kernel.

    Handles device management, kernel launching, and input validation.
    """
    # --- Input Parsing and Validation ---
    probs = kwargs.get("probs", args[0] if len(args) > 0 else None)
    top_k = kwargs.get("top_k", args[1] if len(args) > 1 else None)
    top_p = kwargs.get("top_p", args[2] if len(args) > 2 else None)
    # Allow seed to be passed for deterministic testing
    seed = kwargs.get("seed")

    if probs is None or top_k is None or top_p is None:
        raise ValueError("Inputs 'probs', 'top_k', and 'top_p' must be provided.")

    assert probs.dim() == 2, "probs must be a 2D tensor"
    assert top_k.dim() == 1, "top_k must be a 1D tensor"
    assert top_p.dim() == 1, "top_p must be a 1D tensor"
    
    batch_size, vocab_size = probs.shape
    assert top_k.shape[0] == batch_size, "top_k batch size mismatch"
    assert top_p.shape[0] == batch_size, "top_p batch size mismatch"
    assert vocab_size == 128256, f"vocab_size must be 128256, but got {vocab_size}"

    # --- Device Management ---
    if not torch.cuda.is_available() and probs.device.type != 'cpu':
        raise RuntimeError("This kernel requires a CUDA-enabled GPU.")
    
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    if not torch.cuda.is_available():
        # Fallback to reference for CPU-only environments
        import warnings
        warnings.warn("CUDA not available. Falling back to reference implementation. This will be slow.")
        # Simulating reference run for completeness, as it's not provided
        # In a real scenario, you'd call the actual reference implementation here.
        samples = torch.empty(batch_size, dtype=torch.int64, device='cpu')
        for i in range(batch_size):
            samples[i] = torch.multinomial(probs[i], 1).squeeze()
        return samples

    original_device = probs.device

    # Move all inputs to the GPU where the kernel will run
    probs = probs.to(device, non_blocking=True, dtype=torch.float32)
    top_k = top_k.to(device, non_blocking=True, dtype=torch.int32)
    top_p = top_p.to(device, non_blocking=True, dtype=torch.float32)

    # --- Kernel Launch ---
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    
    # Generate per-row seeds for reproducibility and randomness
    if seed is None:
        seeds = torch.randint(0, 2**31 - 1, (batch_size,), device=device, dtype=torch.int32)
    else:
        # Create a deterministic sequence of seeds if a base seed is provided
        seeds = (torch.arange(batch_size, device=device, dtype=torch.int32) + seed).int()

    grid = (batch_size,)
    
    # Power-of-2 block sizes suitable for tl.sort and modern GPUs.
    # A merge buffer of 1024 (512+512) is efficient for block-level sorting.
    BLOCK_SIZE_M = 512
    TOP_K_BUFFER_SIZE = 512
    
    top_k_top_p_sampling_from_probs_v128256_kernel[grid](
        probs,
        top_k,
        top_p,
        samples,
        seeds,
        batch_size=batch_size,
        stride_probs_b=probs.stride(0),
        VOCAB_SIZE=vocab_size,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        TOP_K_BUFFER_SIZE=TOP_K_BUFFER_SIZE,
    )

    # --- Output Device Management ---
    return samples.to(original_device, non_blocking=True)
scrolls · 246 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON