Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / triton3d9fe1

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def top_p_sampling_kernel(
    probs_ptr,
    top_p_ptr,
    samples_ptr,
    rand_vals_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    # Get batch index
    batch_idx = tl.program_id(0)
    
    if batch_idx >= batch_size:
        return
    
    # Load top_p and random value
    p = tl.load(top_p_ptr + batch_idx)
    random_val = tl.load(rand_vals_ptr + batch_idx)
    
    # Handle p <= 0 case - use argmax
    if p <= 0.0:
        # Find argmax
        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 + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=-1e30
            )
            
            local_max = tl.max(probs, axis=0)
            if local_max > max_val:
                max_val = local_max
                # Find the index of the maximum
                is_max = (probs == local_max) & mask
                indices = tl.where(is_max, block_offsets, vocab_size)
                first_max_idx = tl.min(indices, axis=0)
                if first_max_idx < vocab_size:
                    max_idx = first_max_idx
        
        tl.store(samples_ptr + batch_idx, max_idx)
        return
    
    # For p >= 1.0, sample from full distribution
    if p >= 0.999:  # Effectively p >= 1.0
        cumsum = 0.0
        sampled_idx = 0
        found_sample = 0
        
        for block_start in range(0, vocab_size, BLOCK_SIZE):
            # Skip if already found
            if found_sample == 0:
                block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
                mask = block_offsets < vocab_size
                
                probs = tl.load(
                    probs_ptr + batch_idx * vocab_size + block_offsets,
                    mask=mask,
                    other=0.0
                )
                
                prev_cumsum = cumsum
                block_cumsum = tl.cumsum(probs, axis=0)
                block_cumsum = prev_cumsum + block_cumsum
                
                crosses = (block_cumsum >= random_val) & mask
                if tl.sum(crosses, axis=0) > 0:
                    valid_indices = tl.where(crosses, block_offsets, vocab_size)
                    first_cross = tl.min(valid_indices, axis=0)
                    if first_cross < vocab_size:
                        sampled_idx = first_cross
                        found_sample = 1
                
                cumsum = prev_cumsum + tl.sum(probs, axis=0)
        
        tl.store(samples_ptr + batch_idx, sampled_idx)
        return
    
    # Nucleus sampling: p < 1.0
    # We need to implement top-p filtering
    
    # First pass: find all probabilities and sort them approximately
    # Since we can't sort efficiently in Triton, we'll use a threshold-based approach
    
    # Find the sum of all probabilities (should be ~1.0)
    total_sum = 0.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 + batch_idx * vocab_size + block_offsets,
            mask=mask,
            other=0.0
        )
        
        total_sum += tl.sum(probs, axis=0)
    
    # Binary search for threshold that gives us approximately top-p mass
    low_threshold = 0.0
    high_threshold = 1.0
    
    for _ in range(10):  # 10 iterations of binary search
        mid_threshold = (low_threshold + high_threshold) / 2.0
        
        # Calculate sum of probabilities above threshold
        above_sum = 0.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 + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=0.0
            )
            
            above_mask = (probs >= mid_threshold) & mask
            above_sum += tl.sum(tl.where(above_mask, probs, 0.0), axis=0)
        
        # Adjust threshold based on whether we have too much or too little mass
        if above_sum > p:
            low_threshold = mid_threshold
        else:
            high_threshold = mid_threshold
    
    # Use the final threshold
    threshold = low_threshold
    
    # Calculate the actual sum with this threshold for normalization
    filtered_sum = 0.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 + batch_idx * vocab_size + block_offsets,
            mask=mask,
            other=0.0
        )
        
        above_mask = (probs >= threshold) & mask
        filtered_sum += tl.sum(tl.where(above_mask, probs, 0.0), axis=0)
    
    # Ensure filtered_sum is not zero
    if filtered_sum <= 0.0:
        filtered_sum = 1.0
        threshold = 0.0
    
    # Sample from the filtered distribution
    target = random_val * filtered_sum
    cumsum = 0.0
    sampled_idx = 0
    found_sample = 0
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        # Skip if already found
        if found_sample == 0:
            block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_offsets < vocab_size
            
            probs = tl.load(
                probs_ptr + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=0.0
            )
            
            above_mask = (probs >= threshold) & mask
            filtered_probs = tl.where(above_mask, probs, 0.0)
            
            prev_cumsum = cumsum
            block_cumsum = tl.cumsum(filtered_probs, axis=0)
            block_cumsum = prev_cumsum + block_cumsum
            
            crosses = (block_cumsum >= target) & above_mask
            if tl.sum(crosses, axis=0) > 0:
                valid_indices = tl.where(crosses, block_offsets, vocab_size)
                first_cross = tl.min(valid_indices, axis=0)
                if first_cross < vocab_size:
                    sampled_idx = first_cross
                    found_sample = 1
            
            cumsum = prev_cumsum + tl.sum(filtered_probs, axis=0)
    
    # Fallback: if we didn't find anything, use argmax
    if found_sample == 0:
        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 + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=-1e30
            )
            
            local_max = tl.max(probs, axis=0)
            if local_max > max_val:
                max_val = local_max
                is_max = (probs == local_max) & mask
                indices = tl.where(is_max, block_offsets, vocab_size)
                first_max_idx = tl.min(indices, axis=0)
                if first_max_idx < vocab_size:
                    max_idx = first_max_idx
        
        sampled_idx = max_idx
    
    tl.store(samples_ptr + batch_idx, sampled_idx)


@triton.jit
def top_p_sampling_kernel_simple(
    probs_ptr,
    top_p_ptr,
    samples_ptr,
    rand_vals_ptr,
    batch_size,
    vocab_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Simpler but potentially less accurate nucleus sampling for better performance"""
    batch_idx = tl.program_id(0)
    
    if batch_idx >= batch_size:
        return
    
    p = tl.load(top_p_ptr + batch_idx)
    random_val = tl.load(rand_vals_ptr + batch_idx)
    
    # Handle special cases
    if p <= 0.0:
        # Argmax
        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 + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=-1e30
            )
            
            local_max = tl.max(probs, axis=0)
            if local_max > max_val:
                max_val = local_max
                is_max = (probs == local_max) & mask
                indices = tl.where(is_max, block_offsets, vocab_size)
                first_max_idx = tl.min(indices, axis=0)
                if first_max_idx < vocab_size:
                    max_idx = first_max_idx
        
        tl.store(samples_ptr + batch_idx, max_idx)
        return
    
    # For all other cases, we'll use cumulative sampling
    # with optional filtering based on p value
    
    # If p < 1.0, we use a simple threshold to filter out low probability tokens
    # This is an approximation of true nucleus sampling
    threshold = 0.0
    if p < 0.999:
        # Find max probability
        max_prob = 0.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 + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=0.0
            )
            
            local_max = tl.max(probs, axis=0)
            if local_max > max_prob:
                max_prob = local_max
        
        # Set threshold based on p and max probability
        # Higher p means lower threshold (include more tokens)
        threshold = max_prob * (1.0 - p) * 0.001
    
    # Calculate filtered sum
    filtered_sum = 0.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 + batch_idx * vocab_size + block_offsets,
            mask=mask,
            other=0.0
        )
        
        above_mask = (probs > threshold) & mask
        filtered_sum += tl.sum(tl.where(above_mask, probs, 0.0), axis=0)
    
    # Sample from filtered distribution
    if filtered_sum <= 0.0:
        filtered_sum = 1.0
        threshold = 0.0
    
    target = random_val * filtered_sum
    cumsum = 0.0
    sampled_idx = 0
    found = 0
    
    for block_start in range(0, vocab_size, BLOCK_SIZE):
        if found == 0:
            block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
            mask = block_offsets < vocab_size
            
            probs = tl.load(
                probs_ptr + batch_idx * vocab_size + block_offsets,
                mask=mask,
                other=0.0
            )
            
            above_mask = (probs > threshold) & mask
            filtered_probs = tl.where(above_mask, probs, 0.0)
            
            prev_cumsum = cumsum
            block_cumsum = tl.cumsum(filtered_probs, axis=0)
            block_cumsum = prev_cumsum + block_cumsum
            
            crosses = (block_cumsum >= target) & above_mask
            if tl.sum(crosses, axis=0) > 0:
                valid_indices = tl.where(crosses, block_offsets, vocab_size)
                first_cross = tl.min(valid_indices, axis=0)
                if first_cross < vocab_size:
                    sampled_idx = first_cross
                    found = 1
            
            cumsum = prev_cumsum + tl.sum(filtered_probs, axis=0)
    
    tl.store(samples_ptr + batch_idx, sampled_idx)


def run(*args, **kwargs):
    """Entry point function for top_p_sampling_from_probs_v128256"""
    
    # Handle both positional and keyword arguments
    if len(args) == 2:
        probs, top_p = args
    elif len(args) == 0 and 'probs' in kwargs and 'top_p' in kwargs:
        probs = kwargs['probs']
        top_p = kwargs['top_p']
    else:
        raise ValueError("Expected 2 arguments: probs and top_p")
    
    # 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 != 'cuda':
        probs = probs.cuda()
    if top_p.device.type != 'cuda':
        top_p = top_p.cuda()
    
    # Validate inputs
    batch_size, vocab_size = probs.shape
    assert vocab_size == 128256, f"Expected vocab_size=128256, got {vocab_size}"
    assert probs.dtype == torch.float32, f"Expected probs dtype float32, got {probs.dtype}"
    assert top_p.dtype == torch.float32, f"Expected top_p dtype float32, got {top_p.dtype}"
    assert top_p.shape == (batch_size,), f"Expected top_p shape ({batch_size},), got {top_p.shape}"
    
    # Ensure inputs are contiguous
    if not probs.is_contiguous():
        probs = probs.contiguous()
    if not top_p.is_contiguous():
        top_p = top_p.contiguous()
    
    # Allocate output tensor
    samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
    
    # Pre-generate random values on GPU
    rand_vals = torch.rand(batch_size, dtype=torch.float32, device=probs.device)
    
    # Configure grid
    grid = (batch_size,)
    
    # Optimal block size for B200
    BLOCK_SIZE = 1024
    
    # Use the main kernel which properly handles nucleus sampling
    top_p_sampling_kernel[grid](
        probs,
        top_p,
        samples,
        rand_vals,
        batch_size,
        vocab_size,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    # Move result back to original device if needed
    if probs_device.type != 'cuda':
        samples = samples.cpu()
    
    return samples
scrolls · 420 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON