Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / tritond676e3

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def top_k_sampling_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    rand_vals_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Top-k sampling kernel optimized for B200 GPU."""
    pid = tl.program_id(0)
    
    if pid >= batch_size:
        return
    
    # Load k value and random value for this sequence
    k = tl.load(top_k_ptr + pid)
    rand_val = tl.load(rand_vals_ptr + pid)
    probs_offset = pid * vocab_size
    
    # Initialize output
    sample_idx = 0
    
    # Handle invalid k values - use original distribution
    if k <= 0 or k >= vocab_size:
        # Direct cumulative sum sampling
        cumsum = 0.0
        found_sample = 0
        
        # Process in blocks for better memory access
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            # Load block of probabilities
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            
            # Process each element in the block
            for i in range(BLOCK_SIZE):
                idx = block_start + i
                if idx < vocab_size and found_sample == 0:
                    prob = tl.load(probs_ptr + probs_offset + idx)
                    cumsum += prob
                    if cumsum >= rand_val:
                        sample_idx = idx
                        found_sample = 1
        
    else:
        # Top-k sampling implementation
        # Step 1: Find approximate threshold using heap-like approach
        
        # We'll use multiple passes to find the k-th largest value
        # First pass: find maximum value
        max_val = 0.0
        
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            
            block_probs = tl.load(
                probs_ptr + probs_offset + block_indices,
                mask=mask,
                other=0.0
            )
            
            block_max = tl.max(block_probs, axis=0)
            max_val = tl.maximum(max_val, block_max)
        
        # Binary search for the k-th largest value
        min_val = 0.0
        threshold = max_val
        
        # Perform binary search iterations
        for iter_idx in range(20):  # 20 iterations for good precision
            mid_val = (max_val + min_val) / 2.0
            count = 0
            
            # Count values >= mid_val
            for block_start in range(0, vocab_size, BLOCK_SIZE):
                block_indices = block_start + tl.arange(0, BLOCK_SIZE)
                mask = block_indices < vocab_size
                
                block_probs = tl.load(
                    probs_ptr + probs_offset + block_indices,
                    mask=mask,
                    other=0.0
                )
                
                above_mid = tl.where(mask, block_probs >= mid_val, 0)
                count += tl.sum(above_mid.to(tl.int32), axis=0)
            
            # Adjust search range
            if count > k:
                min_val = mid_val
            else:
                max_val = mid_val
                threshold = mid_val
        
        # Step 2: Compute sum of top-k probabilities
        sum_topk = 0.0
        
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            
            block_probs = tl.load(
                probs_ptr + probs_offset + block_indices,
                mask=mask,
                other=0.0
            )
            
            # Filter to keep only top-k values
            topk_mask = tl.where(mask, block_probs >= threshold, 0)
            filtered_probs = tl.where(topk_mask, block_probs, 0.0)
            sum_topk += tl.sum(filtered_probs, axis=0)
        
        # Avoid division by zero
        sum_topk = tl.maximum(sum_topk, 1e-10)
        
        # Step 3: Sample from renormalized top-k distribution
        target = rand_val * sum_topk
        cumsum = 0.0
        found_sample = 0
        
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            # Process each element in the block
            for i in range(BLOCK_SIZE):
                idx = block_start + i
                if idx < vocab_size and found_sample == 0:
                    prob = tl.load(probs_ptr + probs_offset + idx)
                    if prob >= threshold:
                        cumsum += prob
                        if cumsum >= target:
                            sample_idx = idx
                            found_sample = 1
    
    # Store the sampled index
    tl.store(samples_ptr + pid, sample_idx)


@triton.jit
def top_k_sampling_kernel_fast(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    rand_vals_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Faster version with approximate top-k for large vocabulary."""
    pid = tl.program_id(0)
    
    if pid >= batch_size:
        return
    
    k = tl.load(top_k_ptr + pid)
    rand_val = tl.load(rand_vals_ptr + pid)
    probs_offset = pid * vocab_size
    
    sample_idx = 0
    
    if k <= 0 or k >= vocab_size:
        # Direct sampling from full distribution
        cumsum = 0.0
        
        for idx in range(vocab_size):
            prob = tl.load(probs_ptr + probs_offset + idx)
            cumsum += prob
            if cumsum >= rand_val:
                sample_idx = idx
                tl.store(samples_ptr + pid, sample_idx)
                return
    else:
        # Approximate top-k using histogram-based approach
        # This is faster but slightly less accurate
        
        # Find max value
        max_val = 0.0
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            block_probs = tl.load(
                probs_ptr + probs_offset + block_indices,
                mask=mask,
                other=0.0
            )
            max_val = tl.maximum(max_val, tl.max(block_probs, axis=0))
        
        # Use a simple threshold estimation
        # Start with a high threshold and lower it until we have at least k elements
        threshold = max_val * 0.1  # Start at 10% of max
        
        # Count elements above threshold
        count = 0
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            block_probs = tl.load(
                probs_ptr + probs_offset + block_indices,
                mask=mask,
                other=0.0
            )
            above = tl.where(mask, block_probs >= threshold, 0)
            count += tl.sum(above.to(tl.int32), axis=0)
        
        # Adjust threshold if we don't have enough elements
        if count < k:
            threshold = max_val * 0.01  # Lower threshold
        
        # Compute sum and sample
        sum_topk = 0.0
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_indices = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_indices < vocab_size
            block_probs = tl.load(
                probs_ptr + probs_offset + block_indices,
                mask=mask,
                other=0.0
            )
            filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
            sum_topk += tl.sum(filtered, axis=0)
        
        sum_topk = tl.maximum(sum_topk, 1e-10)
        target = rand_val * sum_topk
        cumsum = 0.0
        
        for idx in range(vocab_size):
            prob = tl.load(probs_ptr + probs_offset + idx)
            if prob >= threshold:
                cumsum += prob
                if cumsum >= target:
                    sample_idx = idx
                    tl.store(samples_ptr + pid, sample_idx)
                    return
    
    tl.store(samples_ptr + pid, sample_idx)


def run(*args, **kwargs):
    """Entry point function for top-k sampling from probabilities."""
    # Handle both args and kwargs
    if len(args) == 2:
        probs, top_k = args
    else:
        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)
    
    if probs is None or top_k is None:
        raise ValueError("Both 'probs' and 'top_k' must be provided")
    
    # Device management
    original_device = probs.device
    original_top_k_device = top_k.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 computation is required")
        probs = probs.cuda()
    
    if not top_k.is_cuda:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available but GPU computation is required")
        top_k = top_k.cuda()
    
    # Validate inputs
    batch_size, vocab_size = probs.shape
    assert vocab_size == 151936, f"vocab_size must be 151936, got {vocab_size}"
    assert top_k.shape[0] == batch_size, "top_k must have same batch size as probs"
    
    # Ensure correct dtypes
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)
    
    # Generate random values for sampling
    rand_vals = torch.rand(batch_size, dtype=torch.float32, device=probs.device)
    
    # Allocate output
    samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
    
    # Configure kernel launch
    # B200 has high memory bandwidth, use larger blocks
    BLOCK_SIZE = 1024
    
    # Launch kernel
    grid = (batch_size,)
    
    # Use the main kernel for accuracy
    top_k_sampling_kernel[grid](
        probs,
        top_k,
        samples,
        rand_vals,
        batch_size,
        vocab_size,
        BLOCK_SIZE,
    )
    
    # Move result back to original device if needed
    if original_device != samples.device:
        samples = samples.to(original_device)
    
    return samples
scrolls · 308 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON