Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / triton7a27f9

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def argmax_kernel(
    probs_ptr,
    samples_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Fast argmax kernel for p <= 0 case"""
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Find argmax across vocabulary
    max_val = -1e30
    max_idx = 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 = tl.load(probs_ptr + pid * vocab_size + block_offsets, mask=mask, other=-1e30)
        
        # Find local maximum
        local_max_val = tl.max(probs)
        
        # If this block contains a new maximum, find its exact position
        if local_max_val > max_val:
            # Check each element in the block
            for i in range(BLOCK_SIZE):
                if block_offsets[i] < vocab_size:
                    if probs[i] == local_max_val:
                        max_val = local_max_val
                        max_idx = block_start + i
                        break
    
    tl.store(samples_ptr + pid, max_idx)


@triton.jit
def full_sampling_kernel(
    probs_ptr,
    samples_ptr,
    rand_vals_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Sampling kernel for p >= 1.0 case (sample from full distribution)"""
    pid = tl.program_id(0)
    if pid >= batch_size:
        return
    
    # Load random value for this sequence
    rand_val = tl.load(rand_vals_ptr + pid)
    
    # Sample using cumulative sum
    cumsum = 0.0
    sampled_idx = vocab_size - 1
    
    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 = tl.load(probs_ptr + pid * vocab_size + block_offsets, mask=mask, other=0.0)
        
        # Check each probability
        for i in range(BLOCK_SIZE):
            if block_offsets[i] < vocab_size:
                cumsum += probs[i]
                if cumsum > rand_val:
                    sampled_idx = block_start + i
                    tl.store(samples_ptr + pid, sampled_idx)
                    return
    
    tl.store(samples_ptr + pid, sampled_idx)


def run(probs, top_p):
    """
    Top-p (nucleus) sampling from probability distributions.
    
    This implementation uses a hybrid approach:
    - Triton kernels for simple cases (argmax when p<=0, full sampling when p>=1)
    - PyTorch for accurate nucleus sampling when 0 < p < 1
    
    Args:
        probs: [batch_size, vocab_size] probability distributions
        top_p: [batch_size] cumulative probability thresholds
    
    Returns:
        samples: [batch_size] sampled token indices
    """
    # 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()
    
    # Ensure both tensors are on the same GPU
    if probs.device != top_p.device:
        top_p = top_p.to(probs.device)
    
    # Validate inputs
    batch_size, vocab_size = probs.shape
    assert vocab_size == 129280, f"Expected vocab_size=129280, got {vocab_size}"
    
    device = probs.device
    probs = probs.to(torch.float32)
    top_p = top_p.to(torch.float32)
    
    # Create output tensor
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    
    # Process each sequence based on its top_p value
    # We'll batch sequences by their sampling strategy
    argmax_mask = top_p <= 0.0
    full_sample_mask = top_p >= 1.0
    nucleus_mask = ~argmax_mask & ~full_sample_mask
    
    # Handle argmax cases with Triton kernel
    argmax_count = argmax_mask.sum().item()
    if argmax_count > 0:
        argmax_indices = torch.where(argmax_mask)[0]
        argmax_probs = probs[argmax_indices]
        argmax_samples = torch.empty(argmax_count, dtype=torch.int64, device=device)
        
        # Launch argmax kernel
        BLOCK_SIZE = 256
        grid = (argmax_count,)
        argmax_kernel[grid](
            argmax_probs,
            argmax_samples,
            argmax_count,
            vocab_size,
            BLOCK_SIZE=BLOCK_SIZE
        )
        
        samples[argmax_indices] = argmax_samples
    
    # Handle full sampling cases with Triton kernel
    full_count = full_sample_mask.sum().item()
    if full_count > 0:
        full_indices = torch.where(full_sample_mask)[0]
        full_probs = probs[full_indices]
        full_samples = torch.empty(full_count, dtype=torch.int64, device=device)
        
        # Generate random values for sampling
        rand_vals = torch.rand(full_count, device=device, dtype=torch.float32)
        
        # Launch full sampling kernel
        BLOCK_SIZE = 256
        grid = (full_count,)
        full_sampling_kernel[grid](
            full_probs,
            full_samples,
            rand_vals,
            full_count,
            vocab_size,
            BLOCK_SIZE=BLOCK_SIZE
        )
        
        samples[full_indices] = full_samples
    
    # Handle nucleus sampling cases with PyTorch (for accuracy)
    nucleus_count = nucleus_mask.sum().item()
    if nucleus_count > 0:
        nucleus_indices = torch.where(nucleus_mask)[0]
        
        # Process each nucleus sampling case
        for idx in nucleus_indices:
            i = idx.item()
            row = probs[i]
            p = float(top_p[i].item())
            
            # Sort probabilities in descending order
            vals, sorted_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 = sorted_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 filtering fails
                filtered = row
            
            # Sample from the filtered distribution
            samples[i] = torch.multinomial(filtered, 1, replacement=True).squeeze(0)
    
    # Move result back to original device if needed
    if probs_device.type == 'cpu':
        samples = samples.cpu()
    
    return samples
scrolls · 220 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON