Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / triton36a928

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-36a928?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

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:d3de2f996c58c8483925bb44adee6d6f172a47e8f5188874de6209754336dc5b
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

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

@triton.jit
def top_p_sampling_kernel_simple(
    probs_ptr, top_p_ptr, samples_ptr, seeds_ptr,
    batch_size, vocab_size,
    BLOCK_SIZE: tl.constexpr
):
    """
    Kernel for simple sampling cases: argmax (p<=0) or full multinomial (p>=1).
    Each thread block handles one sequence in the batch.
    """
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Load top_p value for this sequence
    top_p_val = tl.load(top_p_ptr + pid)
    
    # Base pointer for this sequence's probability distribution
    probs_base = probs_ptr + pid * vocab_size
    
    # Handle degenerate case: top_p <= 0 means argmax
    if top_p_val <= 0.0:
        max_val = -1e30  # Use large negative number instead of inf
        max_idx = 0
        
        # Find argmax over vocabulary
        for i in range(vocab_size):
            val = tl.load(probs_base + i)
            if val > max_val:
                max_val = val
                max_idx = i
        
        tl.store(samples_ptr + pid, max_idx)
        return
    
    # For top_p >= 1.0, do standard multinomial sampling
    # Generate random value
    seed = tl.load(seeds_ptr + pid)
    rand_val = tl.rand(seed, tl.arange(0, 1))[0]
    
    cumsum = 0.0
    for i in range(vocab_size):
        prob = tl.load(probs_base + i)
        cumsum += prob
        if cumsum >= rand_val:
            tl.store(samples_ptr + pid, i)
            return
    
    # Fallback to last token (shouldn't happen with normalized probs)
    tl.store(samples_ptr + pid, vocab_size - 1)


def run(*args, **kwargs):
    """
    Main entry point for top-p sampling from probability distributions.
    
    Args:
        probs: [batch_size, vocab_size] tensor of probabilities (float32)
        top_p: [batch_size] tensor of top-p values (float32)
    
    Returns:
        samples: [batch_size] tensor of sampled token indices (int64)
    """
    # Handle both args and kwargs
    if len(args) >= 2:
        probs, top_p = args[0], args[1]
    else:
        probs = kwargs.get('probs', args[0] if len(args) > 0 else None)
        top_p = kwargs.get('top_p', args[1] if len(args) > 1 else None)
    
    if probs is None or top_p is None:
        raise ValueError("Both 'probs' and 'top_p' tensors are required")
    
    # Check CUDA availability
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
    
    # Store original devices
    probs_device = probs.device
    top_p_device = top_p.device
    
    # Move tensors to GPU if needed
    if probs.device.type == 'cpu':
        probs = probs.cuda()
    if top_p.device.type == 'cpu':
        top_p = top_p.cuda()
    
    # Validate inputs
    assert probs.dim() == 2, f"probs must be 2D, got {probs.dim()}D"
    assert top_p.dim() == 1, f"top_p must be 1D, got {top_p.dim()}D"
    
    batch_size, vocab_size = probs.shape
    assert vocab_size == 151936, f"vocab_size must be 151936, got {vocab_size}"
    assert top_p.shape[0] == batch_size, f"top_p batch size mismatch"
    
    # Ensure correct dtypes
    probs = probs.to(torch.float32)
    top_p = top_p.to(torch.float32)
    
    # Allocate output
    samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
    
    # Process each sequence based on its top_p value
    for i in range(batch_size):
        p = float(top_p[i].item())
        row = probs[i]
        
        if p <= 0.0:
            # Degenerate to argmax
            samples[i] = torch.argmax(row).to(torch.int64)
        elif p < 1.0:
            # Nucleus sampling: keep top tokens until cumulative prob > p
            vals, idx = torch.sort(row, descending=True)
            cdf = torch.cumsum(vals, dim=0)
            
            # Find cutoff: keep tokens until cumulative probability exceeds p
            # Shift mask to keep the first token that crosses p
            to_remove = cdf > p
            to_remove[1:] = to_remove[:-1].clone()
            to_remove[0] = False
            keep = ~to_remove
            keep_idx = idx[keep]
            
            # Build filtered distribution in original index space
            filtered = torch.zeros_like(row)
            filtered[keep_idx] = row[keep_idx]
            
            # Renormalize
            filtered_sum = filtered.sum()
            if filtered_sum > 0:
                filtered = filtered / filtered_sum
            else:
                # Fallback to original distribution if something goes wrong
                filtered = row
            
            # Sample from filtered distribution
            samples[i] = torch.multinomial(filtered, 1, replacement=True).squeeze(0)
        else:
            # p >= 1.0: sample from full distribution
            samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
    
    # Move result back to original device if necessary
    if probs_device.type == 'cpu':
        samples = samples.cpu()
    
    return samples
scrolls · 150 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON