claude-opus-4-1-20250805 / triton7a27f9
claude-opus-4-1-20250805_triton_7a27f9 · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 220 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-7a27f9?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:20643563b6301653acbb4a09a2cd1ed0daf9797819465546b9da75f3417023b5
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py220 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def argmax_kernel(
probs_ptr,
samples_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""Fast argmax kernel for p <= 0 case"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Find argmax across vocabulary
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 + pid * vocab_size + block_offsets, mask=mask, other=-1e30)
# Find local maximum
local_max_val = tl.max(probs)
# If this block contains a new maximum, find its exact position
if local_max_val > max_val:
# Check each element in the block
for i in range(BLOCK_SIZE):
if block_offsets[i] < vocab_size:
if probs[i] == local_max_val:
max_val = local_max_val
max_idx = block_start + i
break
tl.store(samples_ptr + pid, max_idx)
@triton.jit
def full_sampling_kernel(
probs_ptr,
samples_ptr,
rand_vals_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""Sampling kernel for p >= 1.0 case (sample from full distribution)"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load random value for this sequence
rand_val = tl.load(rand_vals_ptr + pid)
# Sample using cumulative sum
cumsum = 0.0
sampled_idx = vocab_size - 1
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 + pid * vocab_size + block_offsets, mask=mask, other=0.0)
# Check each probability
for i in range(BLOCK_SIZE):
if block_offsets[i] < vocab_size:
cumsum += probs[i]
if cumsum > rand_val:
sampled_idx = block_start + i
tl.store(samples_ptr + pid, sampled_idx)
return
tl.store(samples_ptr + pid, sampled_idx)
def run(probs, top_p):
"""
Top-p (nucleus) sampling from probability distributions.
This implementation uses a hybrid approach:
- Triton kernels for simple cases (argmax when p<=0, full sampling when p>=1)
- PyTorch for accurate nucleus sampling when 0 < p < 1
Args:
probs: [batch_size, vocab_size] probability distributions
top_p: [batch_size] cumulative probability thresholds
Returns:
samples: [batch_size] sampled token indices
"""
# 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()
# Ensure both tensors are on the same GPU
if probs.device != top_p.device:
top_p = top_p.to(probs.device)
# Validate inputs
batch_size, vocab_size = probs.shape
assert vocab_size == 129280, f"Expected vocab_size=129280, got {vocab_size}"
device = probs.device
probs = probs.to(torch.float32)
top_p = top_p.to(torch.float32)
# Create output tensor
samples = torch.empty(batch_size, dtype=torch.int64, device=device)
# Process each sequence based on its top_p value
# We'll batch sequences by their sampling strategy
argmax_mask = top_p <= 0.0
full_sample_mask = top_p >= 1.0
nucleus_mask = ~argmax_mask & ~full_sample_mask
# Handle argmax cases with Triton kernel
argmax_count = argmax_mask.sum().item()
if argmax_count > 0:
argmax_indices = torch.where(argmax_mask)[0]
argmax_probs = probs[argmax_indices]
argmax_samples = torch.empty(argmax_count, dtype=torch.int64, device=device)
# Launch argmax kernel
BLOCK_SIZE = 256
grid = (argmax_count,)
argmax_kernel[grid](
argmax_probs,
argmax_samples,
argmax_count,
vocab_size,
BLOCK_SIZE=BLOCK_SIZE
)
samples[argmax_indices] = argmax_samples
# Handle full sampling cases with Triton kernel
full_count = full_sample_mask.sum().item()
if full_count > 0:
full_indices = torch.where(full_sample_mask)[0]
full_probs = probs[full_indices]
full_samples = torch.empty(full_count, dtype=torch.int64, device=device)
# Generate random values for sampling
rand_vals = torch.rand(full_count, device=device, dtype=torch.float32)
# Launch full sampling kernel
BLOCK_SIZE = 256
grid = (full_count,)
full_sampling_kernel[grid](
full_probs,
full_samples,
rand_vals,
full_count,
vocab_size,
BLOCK_SIZE=BLOCK_SIZE
)
samples[full_indices] = full_samples
# Handle nucleus sampling cases with PyTorch (for accuracy)
nucleus_count = nucleus_mask.sum().item()
if nucleus_count > 0:
nucleus_indices = torch.where(nucleus_mask)[0]
# Process each nucleus sampling case
for idx in nucleus_indices:
i = idx.item()
row = probs[i]
p = float(top_p[i].item())
# Sort probabilities in descending order
vals, sorted_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 = sorted_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 filtering fails
filtered = row
# Sample from the filtered distribution
samples[i] = torch.multinomial(filtered, 1, replacement=True).squeeze(0)
# Move result back to original device if needed
if probs_device.type == 'cpu':
samples = samples.cpu()
return samplesscrolls · 220 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON