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