Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805_triton_002913

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def optimized_top_k_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    seed,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Optimized top-k sampling kernel for B200 GPU.
    Uses vectorized operations and efficient memory access patterns.
    """
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Load k value
    k = tl.load(top_k_ptr + pid).to(tl.int32)
    
    # Generate random value for sampling
    offset = pid * 7 + 3
    rand_val = tl.rand(seed, offset)
    
    # Base pointer for this batch
    row_base = probs_ptr + pid * vocab_size
    
    # Check if we need to filter
    need_filter = (k > 0) & (k < vocab_size)
    
    if need_filter:
        # Find threshold using binary search
        # First pass: find range of values
        max_val = 0.0
        
        for start in range(0, vocab_size, BLOCK_SIZE):
            offs = start + tl.arange(0, BLOCK_SIZE)
            mask = offs < vocab_size
            vals = tl.load(row_base + offs, mask=mask, other=0.0)
            max_val = tl.maximum(max_val, tl.max(tl.where(mask, vals, 0.0)))
        
        # Binary search for threshold
        lo = 0.0
        hi = max_val
        
        for _ in range(12):  # More iterations for better precision
            mid = (lo + hi) / 2.0
            cnt = 0
            
            for start in range(0, vocab_size, BLOCK_SIZE):
                offs = start + tl.arange(0, BLOCK_SIZE)
                mask = offs < vocab_size
                vals = tl.load(row_base + offs, mask=mask, other=0.0)
                cnt += tl.sum(((vals > mid) & mask).to(tl.int32))
            
            if cnt >= k:
                lo = mid
            else:
                hi = mid
        
        thresh = lo
        
        # Compute normalization factor
        norm = 0.0
        for start in range(0, vocab_size, BLOCK_SIZE):
            offs = start + tl.arange(0, BLOCK_SIZE)
            mask = offs < vocab_size
            vals = tl.load(row_base + offs, mask=mask, other=0.0)
            keep = (vals > thresh) & mask
            norm += tl.sum(tl.where(keep, vals, 0.0))
        
        if norm <= 0.0:
            norm = 1.0
            thresh = -1.0
        
        # Sample from filtered distribution
        target = rand_val * norm
        acc = 0.0
        result = 0
        found = 0
        
        for start in range(0, vocab_size, BLOCK_SIZE):
            if found == 0:
                offs = start + tl.arange(0, BLOCK_SIZE)
                mask = offs < vocab_size
                vals = tl.load(row_base + offs, mask=mask, other=0.0)
                
                # Filter values
                keep = (vals > thresh) & mask
                vals = tl.where(keep, vals, 0.0)
                
                # Cumulative sum
                cum = tl.cumsum(vals) + acc
                
                # Find first position where cumsum >= target
                hit = (cum >= target) & mask
                
                # Check if we found the target in this block
                has_hit = tl.sum(hit.to(tl.int32)) > 0
                
                if has_hit:
                    # Find the first True position using reduction
                    # Create indices for positions that hit
                    indices = tl.where(hit, offs, vocab_size)
                    # Find minimum index (first hit)
                    min_idx = tl.min(indices)
                    result = min_idx
                    found = 1
                
                acc += tl.sum(vals)
    
    else:
        # Sample from full distribution
        target = rand_val
        acc = 0.0
        result = 0
        found = 0
        
        for start in range(0, vocab_size, BLOCK_SIZE):
            if found == 0:
                offs = start + tl.arange(0, BLOCK_SIZE)
                mask = offs < vocab_size
                vals = tl.load(row_base + offs, mask=mask, other=0.0)
                
                # Cumulative sum
                cum = tl.cumsum(vals) + acc
                
                # Find first position where cumsum >= target
                hit = (cum >= target) & mask
                
                # Check if we found the target in this block
                has_hit = tl.sum(hit.to(tl.int32)) > 0
                
                if has_hit:
                    # Find the first True position using reduction
                    # Create indices for positions that hit
                    indices = tl.where(hit, offs, vocab_size)
                    # Find minimum index (first hit)
                    min_idx = tl.min(indices)
                    result = min_idx
                    found = 1
                
                acc += tl.sum(vals)
    
    tl.store(samples_ptr + pid, result)


@triton.jit
def fallback_top_k_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    seed,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Fallback kernel with simpler logic for debugging.
    """
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Generate random value
    offset = pid * 13 + 7
    rand_val = tl.rand(seed, offset)
    
    # Base pointer for this batch
    row_base = probs_ptr + pid * vocab_size
    
    # Simple sampling without filtering (for debugging)
    target = rand_val
    acc = 0.0
    result = vocab_size - 1  # Default to last token
    
    for start in range(0, vocab_size, BLOCK_SIZE):
        offs = start + tl.arange(0, BLOCK_SIZE)
        mask = offs < vocab_size
        vals = tl.load(row_base + offs, mask=mask, other=0.0)
        
        # Process each value
        for i in range(BLOCK_SIZE):
            if start + i < vocab_size:
                val = tl.sum(tl.where(tl.arange(0, BLOCK_SIZE) == i, vals, 0.0))
                acc += val
                if acc >= target:
                    result = tl.minimum(result, start + i)
    
    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 (float32)
        top_k: [batch_size] number of top tokens to consider (int32)
    
    Returns:
        samples: [batch_size] sampled token indices (int64)
    """
    # Store original devices
    original_probs_device = probs.device
    original_top_k_device = top_k.device
    
    # Move to GPU if needed
    if not probs.is_cuda:
        if torch.cuda.is_available():
            probs = probs.cuda()
        else:
            raise RuntimeError("CUDA is not available but Triton kernel requires GPU")
    
    if not top_k.is_cuda:
        if torch.cuda.is_available():
            top_k = top_k.cuda()
        else:
            raise RuntimeError("CUDA is not available but Triton kernel requires GPU")
    
    # Get dimensions
    batch_size, vocab_size = probs.shape
    
    # Validate dimensions
    assert vocab_size == 128256, f"Expected vocab_size=128256, got {vocab_size}"
    assert top_k.shape == (batch_size,), f"Expected top_k shape ({batch_size},), got {top_k.shape}"
    
    # Convert to required dtypes
    probs = probs.to(torch.float32).contiguous()
    top_k = top_k.to(torch.int32).contiguous()
    
    # Allocate output
    samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
    
    # Generate random seed
    seed = torch.randint(0, 2**31 - 1, (1,), device=probs.device).item()
    
    # Configure grid
    grid = (batch_size,)
    
    # Determine block size based on vocab size
    BLOCK_SIZE = 512  # Good for B200 GPU
    
    # Launch optimized kernel
    try:
        optimized_top_k_kernel[grid](
            probs_ptr=probs,
            top_k_ptr=top_k,
            samples_ptr=samples,
            seed=seed,
            batch_size=batch_size,
            vocab_size=vocab_size,
            BLOCK_SIZE=BLOCK_SIZE,
        )
    except Exception as e:
        # Fallback to simpler kernel if optimization fails
        print(f"Warning: Optimized kernel failed with {e}, using fallback")
        fallback_top_k_kernel[grid](
            probs_ptr=probs,
            top_k_ptr=top_k,
            samples_ptr=samples,
            seed=seed,
            batch_size=batch_size,
            vocab_size=vocab_size,
            BLOCK_SIZE=BLOCK_SIZE,
        )
    
    # Move result back to original device if needed
    if original_probs_device != samples.device:
        samples = samples.to(original_probs_device)
    
    return samples
scrolls · 278 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON