claude-opus-4-1-20250805 / triton36a928
claude-opus-4-1-20250805_triton_36a928 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 150 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-36a928?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:d3de2f996c58c8483925bb44adee6d6f172a47e8f5188874de6209754336dc5b
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py150 lines
import torch
import triton
import triton.language as tl
@triton.jit
def top_p_sampling_kernel_simple(
probs_ptr, top_p_ptr, samples_ptr, seeds_ptr,
batch_size, vocab_size,
BLOCK_SIZE: tl.constexpr
):
"""
Kernel for simple sampling cases: argmax (p<=0) or full multinomial (p>=1).
Each thread block handles one sequence in the batch.
"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load top_p value for this sequence
top_p_val = tl.load(top_p_ptr + pid)
# Base pointer for this sequence's probability distribution
probs_base = probs_ptr + pid * vocab_size
# Handle degenerate case: top_p <= 0 means argmax
if top_p_val <= 0.0:
max_val = -1e30 # Use large negative number instead of inf
max_idx = 0
# Find argmax over vocabulary
for i in range(vocab_size):
val = tl.load(probs_base + i)
if val > max_val:
max_val = val
max_idx = i
tl.store(samples_ptr + pid, max_idx)
return
# For top_p >= 1.0, do standard multinomial sampling
# Generate random value
seed = tl.load(seeds_ptr + pid)
rand_val = tl.rand(seed, tl.arange(0, 1))[0]
cumsum = 0.0
for i in range(vocab_size):
prob = tl.load(probs_base + i)
cumsum += prob
if cumsum >= rand_val:
tl.store(samples_ptr + pid, i)
return
# Fallback to last token (shouldn't happen with normalized probs)
tl.store(samples_ptr + pid, vocab_size - 1)
def run(*args, **kwargs):
"""
Main entry point for top-p sampling from probability distributions.
Args:
probs: [batch_size, vocab_size] tensor of probabilities (float32)
top_p: [batch_size] tensor of top-p values (float32)
Returns:
samples: [batch_size] tensor of sampled token indices (int64)
"""
# Handle both args and kwargs
if len(args) >= 2:
probs, top_p = args[0], args[1]
else:
probs = kwargs.get('probs', args[0] if len(args) > 0 else None)
top_p = kwargs.get('top_p', args[1] if len(args) > 1 else None)
if probs is None or top_p is None:
raise ValueError("Both 'probs' and 'top_p' tensors are required")
# 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 == 'cpu':
probs = probs.cuda()
if top_p.device.type == 'cpu':
top_p = top_p.cuda()
# Validate inputs
assert probs.dim() == 2, f"probs must be 2D, got {probs.dim()}D"
assert top_p.dim() == 1, f"top_p must be 1D, got {top_p.dim()}D"
batch_size, vocab_size = probs.shape
assert vocab_size == 151936, f"vocab_size must be 151936, got {vocab_size}"
assert top_p.shape[0] == batch_size, f"top_p batch size mismatch"
# Ensure correct dtypes
probs = probs.to(torch.float32)
top_p = top_p.to(torch.float32)
# Allocate output
samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
# Process each sequence based on its top_p value
for i in range(batch_size):
p = float(top_p[i].item())
row = probs[i]
if p <= 0.0:
# Degenerate to argmax
samples[i] = torch.argmax(row).to(torch.int64)
elif p < 1.0:
# Nucleus sampling: keep top tokens until cumulative prob > p
vals, idx = torch.sort(row, descending=True)
cdf = torch.cumsum(vals, dim=0)
# Find cutoff: keep tokens until cumulative probability exceeds p
# Shift mask to keep the first token that crosses p
to_remove = cdf > p
to_remove[1:] = to_remove[:-1].clone()
to_remove[0] = False
keep = ~to_remove
keep_idx = idx[keep]
# Build filtered distribution in original index space
filtered = torch.zeros_like(row)
filtered[keep_idx] = row[keep_idx]
# Renormalize
filtered_sum = filtered.sum()
if filtered_sum > 0:
filtered = filtered / filtered_sum
else:
# Fallback to original distribution if something goes wrong
filtered = row
# Sample from filtered distribution
samples[i] = torch.multinomial(filtered, 1, replacement=True).squeeze(0)
else:
# p >= 1.0: sample from full distribution
samples[i] = torch.multinomial(row, 1, replacement=True).squeeze(0)
# Move result back to original device if necessary
if probs_device.type == 'cpu':
samples = samples.cpu()
return samplesscrolls · 150 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON