claude-opus-4-1-20250805 / tritond676e3
claude-opus-4-1-20250805_triton_d676e3 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 308 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-d676e3?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:c9e7bd6677ed111c97cb5b67e46e83a784a2af66f0a51a1b6fe2825b15cbf554
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py308 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def top_k_sampling_kernel(
probs_ptr,
top_k_ptr,
samples_ptr,
rand_vals_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""Top-k sampling kernel optimized for B200 GPU."""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load k value and random value for this sequence
k = tl.load(top_k_ptr + pid)
rand_val = tl.load(rand_vals_ptr + pid)
probs_offset = pid * vocab_size
# Initialize output
sample_idx = 0
# Handle invalid k values - use original distribution
if k <= 0 or k >= vocab_size:
# Direct cumulative sum sampling
cumsum = 0.0
found_sample = 0
# Process in blocks for better memory access
for block_start in range(0, vocab_size, BLOCK_SIZE):
# Load block of probabilities
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
# Process each element in the block
for i in range(BLOCK_SIZE):
idx = block_start + i
if idx < vocab_size and found_sample == 0:
prob = tl.load(probs_ptr + probs_offset + idx)
cumsum += prob
if cumsum >= rand_val:
sample_idx = idx
found_sample = 1
else:
# Top-k sampling implementation
# Step 1: Find approximate threshold using heap-like approach
# We'll use multiple passes to find the k-th largest value
# First pass: find maximum value
max_val = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
block_max = tl.max(block_probs, axis=0)
max_val = tl.maximum(max_val, block_max)
# Binary search for the k-th largest value
min_val = 0.0
threshold = max_val
# Perform binary search iterations
for iter_idx in range(20): # 20 iterations for good precision
mid_val = (max_val + min_val) / 2.0
count = 0
# Count values >= mid_val
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
above_mid = tl.where(mask, block_probs >= mid_val, 0)
count += tl.sum(above_mid.to(tl.int32), axis=0)
# Adjust search range
if count > k:
min_val = mid_val
else:
max_val = mid_val
threshold = mid_val
# Step 2: Compute sum of top-k probabilities
sum_topk = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
# Filter to keep only top-k values
topk_mask = tl.where(mask, block_probs >= threshold, 0)
filtered_probs = tl.where(topk_mask, block_probs, 0.0)
sum_topk += tl.sum(filtered_probs, axis=0)
# Avoid division by zero
sum_topk = tl.maximum(sum_topk, 1e-10)
# Step 3: Sample from renormalized top-k distribution
target = rand_val * sum_topk
cumsum = 0.0
found_sample = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
# Process each element in the block
for i in range(BLOCK_SIZE):
idx = block_start + i
if idx < vocab_size and found_sample == 0:
prob = tl.load(probs_ptr + probs_offset + idx)
if prob >= threshold:
cumsum += prob
if cumsum >= target:
sample_idx = idx
found_sample = 1
# Store the sampled index
tl.store(samples_ptr + pid, sample_idx)
@triton.jit
def top_k_sampling_kernel_fast(
probs_ptr,
top_k_ptr,
samples_ptr,
rand_vals_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""Faster version with approximate top-k for large vocabulary."""
pid = tl.program_id(0)
if pid >= batch_size:
return
k = tl.load(top_k_ptr + pid)
rand_val = tl.load(rand_vals_ptr + pid)
probs_offset = pid * vocab_size
sample_idx = 0
if k <= 0 or k >= vocab_size:
# Direct sampling from full distribution
cumsum = 0.0
for idx in range(vocab_size):
prob = tl.load(probs_ptr + probs_offset + idx)
cumsum += prob
if cumsum >= rand_val:
sample_idx = idx
tl.store(samples_ptr + pid, sample_idx)
return
else:
# Approximate top-k using histogram-based approach
# This is faster but slightly less accurate
# Find max value
max_val = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
max_val = tl.maximum(max_val, tl.max(block_probs, axis=0))
# Use a simple threshold estimation
# Start with a high threshold and lower it until we have at least k elements
threshold = max_val * 0.1 # Start at 10% of max
# Count elements above threshold
count = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
above = tl.where(mask, block_probs >= threshold, 0)
count += tl.sum(above.to(tl.int32), axis=0)
# Adjust threshold if we don't have enough elements
if count < k:
threshold = max_val * 0.01 # Lower threshold
# Compute sum and sample
sum_topk = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_indices = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_indices < vocab_size
block_probs = tl.load(
probs_ptr + probs_offset + block_indices,
mask=mask,
other=0.0
)
filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
sum_topk += tl.sum(filtered, axis=0)
sum_topk = tl.maximum(sum_topk, 1e-10)
target = rand_val * sum_topk
cumsum = 0.0
for idx in range(vocab_size):
prob = tl.load(probs_ptr + probs_offset + idx)
if prob >= threshold:
cumsum += prob
if cumsum >= target:
sample_idx = idx
tl.store(samples_ptr + pid, sample_idx)
return
tl.store(samples_ptr + pid, sample_idx)
def run(*args, **kwargs):
"""Entry point function for top-k sampling from probabilities."""
# Handle both args and kwargs
if len(args) == 2:
probs, top_k = args
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)
if probs is None or top_k is None:
raise ValueError("Both 'probs' and 'top_k' must be provided")
# Device management
original_device = probs.device
original_top_k_device = top_k.device
# Move to GPU if needed
if not probs.is_cuda:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU computation is required")
probs = probs.cuda()
if not top_k.is_cuda:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU computation is required")
top_k = top_k.cuda()
# Validate inputs
batch_size, vocab_size = probs.shape
assert vocab_size == 151936, f"vocab_size must be 151936, got {vocab_size}"
assert top_k.shape[0] == batch_size, "top_k must have same batch size as probs"
# Ensure correct dtypes
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
# Generate random values for sampling
rand_vals = torch.rand(batch_size, dtype=torch.float32, device=probs.device)
# Allocate output
samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
# Configure kernel launch
# B200 has high memory bandwidth, use larger blocks
BLOCK_SIZE = 1024
# Launch kernel
grid = (batch_size,)
# Use the main kernel for accuracy
top_k_sampling_kernel[grid](
probs,
top_k,
samples,
rand_vals,
batch_size,
vocab_size,
BLOCK_SIZE,
)
# Move result back to original device if needed
if original_device != samples.device:
samples = samples.to(original_device)
return samplesscrolls · 308 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON