Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton2c9c7d

gemini-2.5-pro_triton_2c9c7d · gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-2c9c7d?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:e1a74a6add68b5a9f85b0df6980cba14d5dee87c8f9b9dcf494548007188f89c
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 = 4num_warps = 4
shared-memorysmem_combined_packed = tl.zeros((COMBINED_SIZE,), dtype=tl.uint64, scope='shared')

Kernel source

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

# --- Triton Kernel ---

@triton.jit
def _bitonic_sort_step(data_ptr, size, stride, merge_size, ascending):
    """
    Performs one step of a bitonic sort on a 1D array in memory.
    This is designed to be called iteratively to sort an array.
    """
    # Each thread handles one comparison-swap operation.
    # We only need size // 2 threads for this.
    # However, for simplicity in Triton, we launch `size` threads and mask them.
    # A more advanced implementation might use fewer threads.
    idx = tl.program_id(1) * 32 + tl.arange(0, 32)
    
    # Determine which pairs to compare
    group_idx = idx // stride
    inner_idx = idx % stride
    
    # Calculate indices for comparison based on the bitonic sort network structure
    i = group_idx * stride * 2 + inner_idx
    j = i + stride

    # Ensure we are within a merge block of size merge_size
    # The direction of comparison depends on which half of the merge block we are in
    is_upper_half = ((i // merge_size) % 2 == 1)
    
    # Create a mask to avoid out-of-bounds access and redundant computations
    mask = (idx < size // 2)
    
    # Load elements to be compared
    x1 = tl.load(data_ptr + i, mask=mask)
    x2 = tl.load(data_ptr + j, mask=mask)

    # Determine swap condition based on the bitonic sequence and final sort order
    should_swap = (x1 > x2)
    
    # Flip the swap condition based on the desired final sort order and bitonic stage
    if ascending:
        swap_condition = should_swap if not is_upper_half else not should_swap
    else: # descending
        swap_condition = should_swap if is_upper_half else not should_swap

    # Perform conditional swap
    swapped_x1 = tl.where(swap_condition, x2, x1)
    swapped_x2 = tl.where(swap_condition, x1, x2)

    # Store back the swapped elements
    tl.store(data_ptr + i, swapped_x1, mask=mask)
    tl.store(data_ptr + j, swapped_x2, mask=mask)


@triton.jit
def _bitonic_sort_power_of_2(data_ptr, size, ascending):
    """
    Sorts a 1D tl.tensor of a power-of-2 size using a bitonic sorting network.
    `data_ptr` should be a pointer to an array in shared memory.
    This kernel is launched with enough threads to cover the comparisons needed.
    """
    num_stages = tl.static_log2(size)
    for stage in range(num_stages):
        merge_size = 1 << (stage + 1)
        for step in range(stage + 1):
            stride = 1 << (stage - step)
            # This is a conceptual call; the logic is inlined for Triton's JIT.
            # In a real Triton implementation, this would be part of the main kernel loop.
            # For this structure, we assume the sorting logic is called within the kernel.
            # The body of `_bitonic_sort_step` would be here, or called as a utility.
            # Let's assume the logic is inlined for simplicity of the demonstration.
            
            # Inlined _bitonic_sort_step logic for one thread block:
            idx = tl.arange(0, size // 2)
            group_idx = idx // stride
            inner_idx = idx % stride
            i = group_idx * stride * 2 + inner_idx
            j = i + stride
            is_upper_half = ((i // merge_size) % 2 == 1)
            
            x1 = tl.load(data_ptr + i)
            x2 = tl.load(data_ptr + j)
            
            should_swap = (x1 > x2)
            if ascending:
                swap_condition = should_swap if not is_upper_half else not should_swap
            else:
                swap_condition = should_swap if is_upper_half else not should_swap
            
            swapped_x1 = tl.where(swap_condition, x2, x1)
            swapped_x2 = tl.where(swap_condition, x1, x2)
            
            tl.store(data_ptr + i, swapped_x1)
            tl.store(data_ptr + j, swapped_x2)
            tl.sync_threads()


@triton.jit
def _top_k_sampling_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    seed_tensor_ptr,
    stride_probs_b,
    VOCAB_SIZE: tl.constexpr,
    STATIC_K: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    """
    Triton kernel for top-k sampling. Each program instance processes one sequence.
    """
    pid_b = tl.program_id(0)

    # --- Shared Memory Declaration ---
    COMBINED_SIZE = tl.constexpr(STATIC_K + BLOCK_V)
    smem_combined_packed = tl.zeros((COMBINED_SIZE,), dtype=tl.uint64, scope='shared')

    # --- Load `k` and Seed for the current sequence ---
    k = tl.load(top_k_ptr + pid_b)
    seed = tl.load(seed_tensor_ptr + pid_b)

    # --- Conditional execution: Top-K path vs. Full Vocab Path ---
    if (k > 0) and (k < VOCAB_SIZE):
        # --- Top-K Path ---
        top_k_packed = tl.full([STATIC_K], 0, dtype=tl.uint64) # Start with prob=0, idx=0
        v_offsets = tl.arange(0, BLOCK_V)

        for v_start_idx in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
            v_start = v_start_idx * BLOCK_V
            v_range = v_start + v_offsets
            v_mask = v_range < VOCAB_SIZE

            probs = tl.load(probs_ptr + pid_b * stride_probs_b + v_range, mask=v_mask, other=0.0)
            indices = v_range.to(tl.uint32)
            probs_uint32 = tl.view(probs, tl.uint32)
            current_packed = (probs_uint32.to(tl.uint64) << 32) | indices.to(tl.uint64)

            # Merge candidates in shared memory
            tl.store(smem_combined_packed + tl.arange(0, STATIC_K), top_k_packed)
            tl.store(smem_combined_packed + STATIC_K + v_offsets, current_packed, mask=v_mask)
            tl.sync_threads()

            _bitonic_sort_power_of_2(smem_combined_packed, COMBINED_SIZE, ascending=False)
            
            top_k_packed = tl.load(smem_combined_packed + tl.arange(0, STATIC_K))

        # Unpack the final top K candidates
        top_k_indices = (top_k_packed & 0xFFFFFFFF).to(tl.int64)
        top_k_probs = tl.view((top_k_packed >> 32).to(tl.uint32), tl.float32)

        # Gumbel-Max sampling on the filtered top K items
        k_arange = tl.arange(0, STATIC_K)
        k_mask = k_arange < k
        
        log_probs = tl.log(top_k_probs + 1e-9)
        rand_offsets = pid_b * STATIC_K + k_arange
        rand_uniform = tl.rand(seed, rand_offsets)
        gumbel_noise = -tl.log(-tl.log(rand_uniform + 1e-9) + 1e-9)
        gumbel_scores = tl.where(k_mask, log_probs + gumbel_noise, -float('inf'))

        winner_idx_in_block = tl.argmax(gumbel_scores, axis=0)
        sampled_token_id = tl.load(top_k_indices + winner_idx_in_block)

    else:
        # --- Full Vocab Path ---
        max_gumbel_score = -float('inf')
        result_index = -1
        
        v_offsets = tl.arange(0, BLOCK_V)
        for v_start_idx in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
            v_start = v_start_idx * BLOCK_V
            v_range = v_start + v_offsets
            v_mask = v_range < VOCAB_SIZE

            probs = tl.load(probs_ptr + pid_b * stride_probs_b + v_range, mask=v_mask, other=0.0)
            log_probs = tl.log(probs + 1e-9)

            rand_offsets = pid_b * VOCAB_SIZE + v_range
            rand_uniform = tl.rand(seed, rand_offsets)
            gumbel_noise = -tl.log(-tl.log(rand_uniform + 1e-9) + 1e-9)
            gumbel_scores = tl.where(v_mask, log_probs + gumbel_noise, -float('inf'))
            
            block_max_score = tl.max(gumbel_scores, axis=0)
            
            update_mask = block_max_score > max_gumbel_score
            max_gumbel_score = tl.where(update_mask, block_max_score, max_gumbel_score)
            
            block_max_idx = tl.argmax(gumbel_scores, axis=0)
            block_winner_vocab_idx = (v_start + block_max_idx)
            result_index = tl.where(update_mask, block_winner_vocab_idx, result_index)

        sampled_token_id = result_index.to(tl.int64)

    tl.store(samples_ptr + pid_b, sampled_token_id)


def top_k_sampling_from_probs_v129280(probs: torch.Tensor, top_k: torch.Tensor) -> torch.Tensor:
    """
    Performs top-k sampling from probability distributions using a Triton kernel.
    """
    if not torch.cuda.is_available():
        raise RuntimeError("This kernel requires a CUDA-enabled GPU.")
    
    if probs.dim() != 2 or top_k.dim() != 1 or probs.shape[0] != top_k.shape[0]:
        raise ValueError("Invalid shapes. probs must be [batch, vocab], top_k must be [batch].")
    
    batch_size, vocab_size = probs.shape
    assert vocab_size == 129280, "This kernel is specialized for vocab_size=129280"

    # Define kernel constants.
    # Note: Using larger STATIC_K might require more shared memory and register spills,
    # but handles larger k values more efficiently within the fast path.
    # For bitonic sort, (STATIC_K + BLOCK_V) must be a power of 2.
    STATIC_K = 64
    BLOCK_V = 64
    combined_size = STATIC_K + BLOCK_V
    if (combined_size & (combined_size - 1) != 0) or combined_size == 0:
         raise ValueError(f"STATIC_K ({STATIC_K}) + BLOCK_V ({BLOCK_V}) must be a power of two for the bitonic sort.")
    
    original_device = probs.device
    device = torch.device("cuda")

    # Move data to GPU
    probs_gpu = probs.to(device=device, dtype=torch.float32, non_blocking=True)
    top_k_gpu = top_k.to(device=device, dtype=torch.int32, non_blocking=True)

    # Allocate output and seed tensors
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    seed_tensor = torch.randint(0, 2**32 - 1, (batch_size,), dtype=torch.int64, device=device)

    grid = (batch_size,)
    
    # We use one warp per program instance. More complex kernels might need more.
    # The bitonic sort implementation implicitly uses all threads in the block.
    # A single warp (32 threads) is sufficient for vector loads/stores.
    # However, the bitonic sort is most efficient when using more threads.
    # Let's use 4 warps to provide enough parallelism for the sort.
    num_warps = 4
    
    _top_k_sampling_kernel[grid](
        probs_gpu,
        top_k_gpu,
        samples,
        seed_tensor,
        stride_probs_b=probs_gpu.stride(0),
        VOCAB_SIZE=vocab_size,
        STATIC_K=STATIC_K,
        BLOCK_V=BLOCK_V,
        num_warps=num_warps,
    )

    # Move result back to the original device
    return samples.to(device=original_device, non_blocking=True)


def run(*args, **kwargs):
    """
    Public entry point for the kernel.
    Handles device management and positional/keyword arguments.
    """
    if args:
        if len(args) != 2:
            raise ValueError(f"Expected 2 positional arguments (probs, top_k), but got {len(args)}")
        probs, top_k = args
    elif kwargs:
        try:
            probs = kwargs['probs']
            top_k = kwargs['top_k']
        except KeyError as e:
            raise KeyError(f"Missing required keyword argument: {e}")
    else:
        raise ValueError("No arguments provided. Please provide 'probs' and 'top_k'.")

    return top_k_sampling_from_probs_v129280(probs, top_k)
scrolls · 276 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON