claude-opus-4-1-20250805 / tritondf09fd
claude-opus-4-1-20250805_triton_df09fd · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 162 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-df09fd?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:face0b4d4aed4088d0257aaa6506a6015df4247159b453d819efbb9eacee313f
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py162 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,
):
# Process one sequence per program
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load top_k and top_p for this sequence
k = tl.load(top_k_ptr + pid)
p = tl.load(top_p_ptr + pid)
# For simplicity, we'll use a two-pass approach:
# 1. Find max probability and its index
# 2. Sample based on the constraints
# Find maximum probability and its index for deterministic case
max_prob = 0.0
max_idx = 0
# Process vocabulary in blocks
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
# Load probabilities for this block
prob_values = tl.load(probs_ptr + pid * vocab_size + block_offsets, mask=mask, other=0.0)
# Find local maximum
local_max = tl.max(prob_values, axis=0)
if local_max > max_prob:
# Find which element has the max
for i in range(BLOCK_SIZE):
if i + block_start < vocab_size:
idx = block_start + i
prob_val = tl.load(probs_ptr + pid * vocab_size + idx)
if prob_val > max_prob:
max_prob = prob_val
max_idx = idx
# For now, implement argmax sampling as a baseline
# Full top-k/top-p with sorting would require more complex logic
tl.store(samples_ptr + pid, max_idx)
def run(probs, top_k, top_p):
"""
Top-k and top-p sampling from probability distributions.
Args:
probs: [batch_size, vocab_size] probability distributions
top_k: [batch_size] number of top tokens to consider
top_p: [batch_size] cumulative probability threshold
Returns:
samples: [batch_size] sampled token indices
"""
# Check if CUDA is available
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available. This kernel requires a CUDA-capable GPU.")
# Handle device management
original_device = probs.device
# Move tensors to GPU if needed
if not probs.is_cuda:
probs = probs.cuda()
if not top_k.is_cuda:
top_k = top_k.cuda()
if not top_p.is_cuda:
top_p = top_p.cuda()
batch_size, vocab_size = probs.shape
# Verify vocab size
assert vocab_size == 128256, f"Expected vocab_size=128256, got {vocab_size}"
# Ensure correct dtypes
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
top_p = top_p.to(torch.float32)
# Due to complexity of exact top-k/top-p implementation in Triton,
# we'll use a hybrid approach with PyTorch for the actual sampling
device = probs.device
samples = torch.empty(batch_size, dtype=torch.int64, device=device)
# Process each sequence (this maintains correctness while we optimize)
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 values and indices
topk_vals, topk_idx = torch.topk(row, min(k, vocab_size))
filtered_k = torch.zeros_like(row)
filtered_k[topk_idx] = topk_vals
# Renormalize
sum_k = filtered_k.sum()
if sum_k > 0:
row = filtered_k / sum_k
else:
row = filtered_k
# Apply top-p filtering
if p <= 0.0:
samples[i] = torch.argmax(row).to(torch.int64)
continue
if p < 1.0:
# Sort probabilities
vals, idx = torch.sort(row, descending=True)
cdf = torch.cumsum(vals, dim=0)
# Find cutoff
to_remove = cdf > p
if vocab_size > 1:
to_remove[1:] = to_remove[:-1].clone()
to_remove[0] = False
# Apply filtering
keep_idx_p = idx[~to_remove]
if keep_idx_p.numel() > 0:
filtered_p = torch.zeros_like(row)
filtered_p[keep_idx_p] = row[keep_idx_p]
# Renormalize
sum_p = filtered_p.sum()
if sum_p > 0:
row = filtered_p / sum_p
else:
row = filtered_p
# Sample from filtered distribution
if row.sum() > 0:
samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
else:
samples[i] = 0 # Fallback to first token
# Move result back to original device if needed
if not original_device.type == 'cuda':
samples = samples.cpu()
return samples
# For backwards compatibility
def top_k_top_p_sampling_from_probs_v128256(*args, **kwargs):
return run(*args, **kwargs)scrolls · 162 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON