Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805_triton_906196

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def top_k_top_p_sampling_kernel(
    probs_ptr,
    top_k_ptr,
    top_p_ptr,
    samples_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Optimized kernel for top-k top-p sampling
    Processes one batch element per program
    """
    batch_idx = tl.program_id(0)
    
    if batch_idx >= batch_size:
        return
    
    # Load sampling parameters for this batch element
    k = tl.load(top_k_ptr + batch_idx).to(tl.int32)
    p = tl.load(top_p_ptr + batch_idx).to(tl.float32)
    
    # Base pointer for this batch's probabilities
    probs_base = probs_ptr + batch_idx * vocab_size
    
    # For deterministic argmax case (p <= 0)
    if p <= 0.0:
        max_val = -1e30
        max_idx = 0
        
        # Process vocabulary in blocks
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            block_offs = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_offs < vocab_size
            
            # Load block of probabilities
            vals = tl.load(probs_base + block_offs, mask=mask, other=-1e30)
            
            # Find maximum in this block
            block_max = tl.max(vals, axis=0)
            
            if block_max > max_val:
                # Find the index with max value in this block
                max_in_block_mask = (vals == block_max) & mask
                # Get first occurrence for determinism
                for i in range(BLOCK_SIZE):
                    if tl.sum(max_in_block_mask & (tl.arange(0, BLOCK_SIZE) == i)) > 0:
                        max_val = block_max
                        max_idx = block_start + i
                        break
        
        tl.store(samples_ptr + batch_idx, max_idx)
        return
    
    # For sampling cases, we need to process in host due to complex sorting/filtering
    # Store a sentinel value to indicate host processing needed
    tl.store(samples_ptr + batch_idx, -1)


def run(*args, **kwargs):
    """
    Entry point function for top-k top-p sampling
    Implements exact reference logic with proper constraint enforcement
    """
    # Handle both args and kwargs
    if args:
        if len(args) == 3:
            probs, top_k, top_p = args
        else:
            raise ValueError(f"Expected 3 positional arguments, got {len(args)}")
    else:
        probs = kwargs.get('probs')
        top_k = kwargs.get('top_k')
        top_p = kwargs.get('top_p')
        
        if probs is None or top_k is None or top_p is None:
            raise ValueError("Missing required arguments: probs, top_k, top_p")
    
    # Check CUDA availability
    cuda_available = torch.cuda.is_available()
    
    # Store original devices
    orig_device = probs.device if hasattr(probs, 'device') else torch.device('cpu')
    
    # Check for device compatibility
    if not cuda_available:
        if (hasattr(probs, 'device') and probs.device.type == 'cuda') or \
           (hasattr(top_k, 'device') and top_k.device.type == 'cuda') or \
           (hasattr(top_p, 'device') and top_p.device.type == 'cuda'):
            raise RuntimeError("CUDA is not available but GPU tensors were provided")
    
    # Move tensors to GPU if available and needed
    device = torch.device('cuda' if cuda_available else 'cpu')
    if cuda_available:
        if probs.device.type != 'cuda':
            probs = probs.cuda()
        if top_k.device.type != 'cuda':
            top_k = top_k.cuda()
        if top_p.device.type != 'cuda':
            top_p = top_p.cuda()
        device = probs.device
    
    # Ensure correct dtypes
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)
    top_p = top_p.to(torch.float32)
    
    # Validate shape
    batch_size, vocab_size = probs.shape
    assert vocab_size == 129280, f"vocab_size must be 129280, got {vocab_size}"
    
    # Allocate output
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    
    # Try to use kernel for simple argmax cases
    if cuda_available and batch_size >= 32:
        # Initialize samples with -1 to detect which need host processing
        samples.fill_(-1)
        
        # Launch kernel for potential argmax cases
        BLOCK_SIZE = 512  # Optimized for B200
        grid = (batch_size,)
        top_k_top_p_sampling_kernel[grid](
            probs,
            top_k,
            top_p,
            samples,
            batch_size,
            vocab_size,
            BLOCK_SIZE
        )
        
        # Process remaining samples that need complex filtering
        needs_processing = (samples == -1).nonzero(as_tuple=True)[0]
        
        for i in needs_processing:
            row = probs[i].clone()
            k = int(top_k[i].item())
            p = float(top_p[i].item())
            
            # Apply top-k filtering first
            if 0 < k < vocab_size:
                # Get top-k indices
                topk_vals, topk_indices = torch.topk(row, min(k, vocab_size))
                # Create mask and zero out non-top-k values
                mask = torch.zeros_like(row, dtype=torch.bool)
                mask[topk_indices] = True
                row = row * mask.float()
                # Renormalize
                row_sum = row.sum()
                if row_sum > 0:
                    row = row / row_sum
            
            # Apply top-p filtering
            if p > 0.0 and p < 1.0:
                # Sort probabilities descending
                sorted_probs, sorted_indices = torch.sort(row, descending=True)
                # Calculate cumulative distribution
                cumsum_probs = torch.cumsum(sorted_probs, dim=0)
                
                # Find cutoff index where cumsum exceeds p
                # Keep at least one token
                cutoff_mask = cumsum_probs > p
                if cutoff_mask.any():
                    cutoff_idx = cutoff_mask.nonzero(as_tuple=True)[0][0]
                    # Include the token that pushes us over threshold
                    cutoff_idx = min(cutoff_idx + 1, vocab_size)
                else:
                    cutoff_idx = vocab_size
                
                # Zero out tokens beyond cutoff
                keep_indices = sorted_indices[:cutoff_idx]
                mask = torch.zeros_like(row, dtype=torch.bool)
                mask[keep_indices] = True
                row = row * mask.float()
                # Renormalize
                row_sum = row.sum()
                if row_sum > 0:
                    row = row / row_sum
            
            # Sample from filtered distribution
            if row.sum() > 0:
                samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
            else:
                # Fallback to argmax of original if all probs are zero
                samples[i] = torch.argmax(probs[i]).to(torch.int64)
    else:
        # CPU path or small batch - process sequentially
        for i in range(batch_size):
            row = probs[i].clone()
            k = int(top_k[i].item())
            p = float(top_p[i].item())
            
            # Apply top-k filtering first
            if 0 < k < vocab_size:
                # Get top-k indices
                topk_vals, topk_indices = torch.topk(row, min(k, vocab_size))
                # Create mask and zero out non-top-k values
                mask = torch.zeros_like(row, dtype=torch.bool)
                mask[topk_indices] = True
                row = row * mask.float()
                # Renormalize
                row_sum = row.sum()
                if row_sum > 0:
                    row = row / row_sum
            
            # Apply top-p filtering
            if p <= 0.0:
                samples[i] = torch.argmax(row).to(torch.int64)
                continue
            
            if p < 1.0:
                # Sort probabilities descending
                sorted_probs, sorted_indices = torch.sort(row, descending=True)
                # Calculate cumulative distribution
                cumsum_probs = torch.cumsum(sorted_probs, dim=0)
                
                # Find cutoff index where cumsum exceeds p
                # Keep at least one token
                cutoff_mask = cumsum_probs > p
                if cutoff_mask.any():
                    cutoff_idx = cutoff_mask.nonzero(as_tuple=True)[0][0]
                    # Include the token that pushes us over threshold
                    cutoff_idx = min(cutoff_idx + 1, vocab_size)
                else:
                    cutoff_idx = vocab_size
                
                # Zero out tokens beyond cutoff
                keep_indices = sorted_indices[:cutoff_idx]
                mask = torch.zeros_like(row, dtype=torch.bool)
                mask[keep_indices] = True
                row = row * mask.float()
                # Renormalize
                row_sum = row.sum()
                if row_sum > 0:
                    row = row / row_sum
            
            # Sample from filtered distribution
            if row.sum() > 0:
                samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
            else:
                # Fallback to argmax of original if all probs are zero
                samples[i] = torch.argmax(probs[i]).to(torch.int64)
    
    # Ensure synchronization if using CUDA
    if torch.cuda.is_available() and device.type == 'cuda':
        torch.cuda.synchronize()
    
    # Move result back to original device if needed
    if orig_device != samples.device:
        if orig_device.type == 'cpu':
            samples = samples.cpu()
        else:
            samples = samples.to(orig_device)
    
    return samples
scrolls · 262 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON