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 samplesscrolls · 262 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON