Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / tritonafd42d

claude-opus-4-1-20250805_triton_afd42d · claude-opus-4-1-20250805 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-afd42d?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:0625c161771b72efef988a487d1ee57bac17a0f822f66051c437ad7c8378fa00
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

main.py315 lines
import torch
import triton
import triton.language as tl

@triton.jit
def top_k_sampling_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    seeds_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Each program handles one sequence
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Load k for this sequence
    k = tl.load(top_k_ptr + pid).to(tl.int32)
    
    # Load random seed for this sequence
    seed = tl.load(seeds_ptr + pid)
    
    # If k is invalid, sample from full distribution
    if k <= 0 or k >= vocab_size:
        k = vocab_size
    
    # We'll do multiple passes to find top-k values
    # First pass: find maximum
    max_val = 0.0
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        block_max = tl.max(block_probs, axis=0)
        max_val = tl.maximum(max_val, block_max)
    
    # Binary search for threshold that gives us exactly k elements
    # We'll find the k-th largest value
    low = 0.0
    high = max_val
    threshold = max_val
    
    for _ in range(20):  # 20 iterations should be enough for convergence
        mid = (low + high) / 2.0
        
        # Count how many elements are >= mid
        count = 0
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_offsets < vocab_size
            
            probs_offsets = pid * vocab_size + block_offsets
            block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
            
            above_threshold = (block_probs >= mid).to(tl.int32)
            count += tl.sum(above_threshold, axis=0)
        
        if count > k:
            low = mid
        else:
            high = mid
            threshold = mid
    
    # Now compute sum of top-k probabilities for renormalization
    sum_topk = 0.0
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        # Only keep probabilities >= threshold
        keep_mask = block_probs >= threshold
        filtered_probs = tl.where(keep_mask, block_probs, 0.0)
        sum_topk += tl.sum(filtered_probs, axis=0)
    
    # Prevent division by zero
    if sum_topk <= 0.0:
        sum_topk = 1.0
    
    # Generate random number for sampling
    random_offset = pid * 4 + tl.arange(0, 1)
    random_val = tl.rand(seed, random_offset) * sum_topk
    
    # Perform sampling by accumulating probabilities
    cumsum = 0.0
    sampled_idx = 0
    found = 0
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        # Only keep probabilities >= threshold
        keep_mask = block_probs >= threshold
        filtered_probs = tl.where(keep_mask, block_probs, 0.0)
        
        # Compute cumulative sum for this block
        # We need to process elements sequentially for cumsum
        # Use a reduction approach instead
        prev_cumsum = cumsum
        block_cumsum = tl.cumsum(filtered_probs, axis=0) + prev_cumsum
        cumsum = prev_cumsum + tl.sum(filtered_probs, axis=0)
        
        # Check if random value falls in this block
        sample_mask = (block_cumsum >= random_val) & (filtered_probs > 0) & (found == 0)
        
        # Find the first position where sample_mask is true
        # We'll use a reduction to find the minimum index where condition is true
        indices_where_true = tl.where(sample_mask, block_offsets, vocab_size)
        min_idx = tl.min(indices_where_true, axis=0)
        
        if min_idx < vocab_size:
            sampled_idx = min_idx
            found = 1
    
    # Fallback: if no sample was found (shouldn't happen), sample the first valid token
    if found == 0:
        sampled_idx = 0
    
    # Store the sampled index
    tl.store(samples_ptr + pid, sampled_idx)


@triton.jit
def top_k_sampling_kernel_simple(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    seeds_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Simplified version that's more robust"""
    # Each program handles one sequence
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Load k for this sequence
    k = tl.load(top_k_ptr + pid).to(tl.int32)
    
    # Load random seed for this sequence
    seed = tl.load(seeds_ptr + pid)
    
    # If k is invalid, sample from full distribution
    if k <= 0 or k >= vocab_size:
        k = vocab_size
    
    # Find the k-th largest value using sorting approach
    # We'll use a simpler approach: find threshold iteratively
    
    # First, find min and max values
    min_val = 1.0
    max_val = 0.0
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        block_max = tl.max(tl.where(mask, block_probs, 0.0), axis=0)
        block_min = tl.min(tl.where(mask & (block_probs > 0), block_probs, 1.0), axis=0)
        
        max_val = tl.maximum(max_val, block_max)
        min_val = tl.minimum(min_val, block_min)
    
    # Binary search for the k-th largest value
    threshold = min_val
    for _ in range(30):  # More iterations for better precision
        mid = (min_val + max_val) / 2.0
        
        # Count elements >= mid
        count = 0
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_offsets < vocab_size
            
            probs_offsets = pid * vocab_size + block_offsets
            block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
            
            above = (block_probs >= mid).to(tl.int32)
            count += tl.sum(above, axis=0)
        
        if count > k:
            min_val = mid
        else:
            max_val = mid
            threshold = mid
    
    # Compute sum for renormalization
    sum_topk = 0.0
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
        sum_topk += tl.sum(filtered, axis=0)
    
    # Generate random value
    random_offset = pid
    rand_val = tl.rand(seed, random_offset + tl.arange(0, 1)) * sum_topk
    rand_scalar = tl.sum(rand_val, axis=0)  # Convert to scalar
    
    # Sample using cumsum
    cumsum = 0.0
    result = vocab_size - 1  # Default to last token
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
        mask = block_offsets < vocab_size
        
        probs_offsets = pid * vocab_size + block_offsets
        block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
        
        filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
        
        # Check each position using vectorized operations
        prev_cumsum = cumsum
        cumsum_vec = tl.cumsum(filtered, axis=0) + prev_cumsum
        cumsum = prev_cumsum + tl.sum(filtered, axis=0)
        
        # Find first position where cumsum >= random
        above_random = cumsum_vec >= rand_scalar
        valid = above_random & mask & (filtered > 0)
        
        # Get minimum index where condition is true
        indices = tl.where(valid, block_offsets, vocab_size)
        min_idx = tl.min(indices, axis=0)
        
        # Update result if we found a valid index
        if min_idx < vocab_size and min_idx < result:
            result = min_idx
    
    # Store result
    tl.store(samples_ptr + pid, result)


def run(probs, top_k):
    """
    Top-k sampling from probability distributions.
    
    Args:
        probs: [batch_size, vocab_size] probability distributions
        top_k: [batch_size] number of top tokens to consider
    
    Returns:
        samples: [batch_size] sampled token indices
    """
    # Store original device
    original_device = probs.device
    
    # Move to GPU if needed
    if not probs.is_cuda:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU operation is required")
        probs = probs.cuda()
    
    if not top_k.is_cuda:
        top_k = top_k.cuda() if torch.cuda.is_available() else top_k
        if not top_k.is_cuda:
            raise RuntimeError("CUDA is not available but GPU operation is required")
    
    # Validate inputs
    batch_size, vocab_size = probs.shape
    assert vocab_size == 129280, f"Expected vocab_size=129280, got {vocab_size}"
    assert top_k.shape == (batch_size,), f"top_k shape mismatch"
    
    # Convert to required dtypes
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)
    
    # Allocate output
    samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
    
    # Generate random seeds
    seeds = torch.randint(0, 2**31-1, (batch_size,), dtype=torch.int32, device=probs.device)
    
    # Choose block size - 512 works well for this vocab size
    BLOCK_SIZE = 512
    
    # Launch kernel
    grid = (batch_size,)
    top_k_sampling_kernel_simple[grid](
        probs,
        top_k,
        samples,
        seeds,
        batch_size,
        vocab_size,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    # Move result back to original device if needed
    if original_device.type != 'cuda':
        samples = samples.cpu()
    
    return samples
scrolls · 315 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON