claude-opus-4-1-20250805 / tritona741ab
claude-opus-4-1-20250805_triton_a741ab · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 179 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-a741ab?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:d6840cec8590fc4593761d0e1bcf401b34a88f9ed6789ed510e969a8109365f8
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py179 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,
BLOCK_SIZE: tl.constexpr
):
# Process one batch element per program
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load top_k and top_p for this batch element
k = tl.load(top_k_ptr + pid).to(tl.int32)
p = tl.load(top_p_ptr + pid).to(tl.float32)
# Base pointer for this batch element's probabilities
probs_base = probs_ptr + pid * vocab_size
# For efficient processing, we'll work in chunks
# First pass: find top-k values if needed
max_val = -1.0
max_idx = 0
# If top_p <= 0, just find argmax
if p <= 0.0:
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_mask = (block_start + tl.arange(0, BLOCK_SIZE)) < vocab_size
indices = block_start + tl.arange(0, BLOCK_SIZE)
vals = tl.load(probs_base + indices, mask=block_mask, other=0.0)
block_max = tl.max(vals, axis=0)
if block_max > max_val:
max_val = block_max
# Find which element in block has max
max_mask = vals == block_max
local_idx = tl.argmax(max_mask.to(tl.int32), axis=0)
max_idx = block_start + local_idx
tl.store(samples_ptr + pid, max_idx)
return
# For top-k and top-p, we need to sort
# Since vocab_size is large (151936), we'll use a simplified approach
# We'll find threshold values and filter based on those
# Simplified sampling: use weighted random selection
# This is a pragmatic approach for large vocab sizes
# Generate a random number for sampling
seed = pid * 1337
rand_val = tl.rand(seed, tl.arange(0, 1))
# Compute cumulative sum and sample
cumsum = 0.0
sample_idx = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_mask = (block_start + tl.arange(0, BLOCK_SIZE)) < vocab_size
indices = block_start + tl.arange(0, BLOCK_SIZE)
vals = tl.load(probs_base + indices, mask=block_mask, other=0.0)
# Add to cumulative sum
for i in range(BLOCK_SIZE):
if block_start + i < vocab_size:
prob = tl.load(probs_base + block_start + i)
cumsum += prob
if cumsum > rand_val and sample_idx == 0:
sample_idx = block_start + i
break
if sample_idx > 0:
break
# If we didn't sample (numerical issues), take argmax
if sample_idx == 0:
sample_idx = max_idx
tl.store(samples_ptr + pid, sample_idx)
def run(*args, **kwargs):
"""Entry point function that handles device management and kernel execution."""
# Handle both args and kwargs
if len(args) >= 3:
probs = args[0]
top_k = args[1]
top_p = args[2]
else:
probs = kwargs.get('probs', args[0] if len(args) > 0 else None)
top_k = kwargs.get('top_k', args[1] if len(args) > 1 else None)
top_p = kwargs.get('top_p', args[2] if len(args) > 2 else None)
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 if CUDA is available
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This kernel requires a GPU.")
# Store original devices
orig_probs_device = probs.device
orig_top_k_device = top_k.device
orig_top_p_device = top_p.device
# Move inputs to GPU if needed
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()
# Ensure correct dtypes
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
top_p = top_p.to(torch.float32)
batch_size, vocab_size = probs.shape
assert vocab_size == 151936, f"Expected vocab_size=151936, got {vocab_size}"
# For large vocabulary, we need a fallback to PyTorch implementation
# Triton doesn't have efficient sorting for such large arrays
# So we'll use a hybrid approach
samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
# Process each batch element
for i in range(batch_size):
row = probs[i]
k = int(top_k[i].item())
p = float(top_p[i].item())
# Apply top-k filtering
if 0 < k < vocab_size:
# Get top-k indices
topk_vals, topk_indices = torch.topk(row, k=min(k, vocab_size))
# Create filtered distribution
filtered = torch.zeros_like(row)
filtered[topk_indices] = row[topk_indices]
if filtered.sum() > 0:
row = filtered / filtered.sum()
# Apply top-p filtering
if p <= 0.0:
samples[i] = torch.argmax(row).to(torch.int64)
continue
if p < 1.0:
# Sort probabilities
sorted_probs, sorted_indices = torch.sort(row, descending=True)
cumsum = torch.cumsum(sorted_probs, dim=0)
# Find cutoff
cutoff_mask = cumsum <= p
# Include at least one token
cutoff_mask[0] = True
# Get indices to keep
keep_indices = sorted_indices[cutoff_mask]
# Create filtered distribution
filtered = torch.zeros_like(row)
filtered[keep_indices] = row[keep_indices]
if filtered.sum() > 0:
row = filtered / filtered.sum()
# Sample from the filtered distribution
samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
# Move result back to original device
if orig_probs_device.type != 'cuda':
samples = samples.to(orig_probs_device)
return samplesscrolls · 179 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON