claude-opus-4-1-20250805_triton_002913
claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 278 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-002913?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:cae1d057334e3cc31df36edd881d7f44682f6cde7bb9acdc06c6ea8ceb81ab9a
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py278 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def optimized_top_k_kernel(
probs_ptr,
top_k_ptr,
samples_ptr,
seed,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""
Optimized top-k sampling kernel for B200 GPU.
Uses vectorized operations and efficient memory access patterns.
"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load k value
k = tl.load(top_k_ptr + pid).to(tl.int32)
# Generate random value for sampling
offset = pid * 7 + 3
rand_val = tl.rand(seed, offset)
# Base pointer for this batch
row_base = probs_ptr + pid * vocab_size
# Check if we need to filter
need_filter = (k > 0) & (k < vocab_size)
if need_filter:
# Find threshold using binary search
# First pass: find range of values
max_val = 0.0
for start in range(0, vocab_size, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
max_val = tl.maximum(max_val, tl.max(tl.where(mask, vals, 0.0)))
# Binary search for threshold
lo = 0.0
hi = max_val
for _ in range(12): # More iterations for better precision
mid = (lo + hi) / 2.0
cnt = 0
for start in range(0, vocab_size, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
cnt += tl.sum(((vals > mid) & mask).to(tl.int32))
if cnt >= k:
lo = mid
else:
hi = mid
thresh = lo
# Compute normalization factor
norm = 0.0
for start in range(0, vocab_size, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
keep = (vals > thresh) & mask
norm += tl.sum(tl.where(keep, vals, 0.0))
if norm <= 0.0:
norm = 1.0
thresh = -1.0
# Sample from filtered distribution
target = rand_val * norm
acc = 0.0
result = 0
found = 0
for start in range(0, vocab_size, BLOCK_SIZE):
if found == 0:
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
# Filter values
keep = (vals > thresh) & mask
vals = tl.where(keep, vals, 0.0)
# Cumulative sum
cum = tl.cumsum(vals) + acc
# Find first position where cumsum >= target
hit = (cum >= target) & mask
# Check if we found the target in this block
has_hit = tl.sum(hit.to(tl.int32)) > 0
if has_hit:
# Find the first True position using reduction
# Create indices for positions that hit
indices = tl.where(hit, offs, vocab_size)
# Find minimum index (first hit)
min_idx = tl.min(indices)
result = min_idx
found = 1
acc += tl.sum(vals)
else:
# Sample from full distribution
target = rand_val
acc = 0.0
result = 0
found = 0
for start in range(0, vocab_size, BLOCK_SIZE):
if found == 0:
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
# Cumulative sum
cum = tl.cumsum(vals) + acc
# Find first position where cumsum >= target
hit = (cum >= target) & mask
# Check if we found the target in this block
has_hit = tl.sum(hit.to(tl.int32)) > 0
if has_hit:
# Find the first True position using reduction
# Create indices for positions that hit
indices = tl.where(hit, offs, vocab_size)
# Find minimum index (first hit)
min_idx = tl.min(indices)
result = min_idx
found = 1
acc += tl.sum(vals)
tl.store(samples_ptr + pid, result)
@triton.jit
def fallback_top_k_kernel(
probs_ptr,
top_k_ptr,
samples_ptr,
seed,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""
Fallback kernel with simpler logic for debugging.
"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# Generate random value
offset = pid * 13 + 7
rand_val = tl.rand(seed, offset)
# Base pointer for this batch
row_base = probs_ptr + pid * vocab_size
# Simple sampling without filtering (for debugging)
target = rand_val
acc = 0.0
result = vocab_size - 1 # Default to last token
for start in range(0, vocab_size, BLOCK_SIZE):
offs = start + tl.arange(0, BLOCK_SIZE)
mask = offs < vocab_size
vals = tl.load(row_base + offs, mask=mask, other=0.0)
# Process each value
for i in range(BLOCK_SIZE):
if start + i < vocab_size:
val = tl.sum(tl.where(tl.arange(0, BLOCK_SIZE) == i, vals, 0.0))
acc += val
if acc >= target:
result = tl.minimum(result, start + i)
tl.store(samples_ptr + pid, result)
def run(probs, top_k):
"""
Top-k sampling from probability distributions.
Args:
probs: [batch_size, vocab_size] probability distributions (float32)
top_k: [batch_size] number of top tokens to consider (int32)
Returns:
samples: [batch_size] sampled token indices (int64)
"""
# Store original devices
original_probs_device = probs.device
original_top_k_device = top_k.device
# Move to GPU if needed
if not probs.is_cuda:
if torch.cuda.is_available():
probs = probs.cuda()
else:
raise RuntimeError("CUDA is not available but Triton kernel requires GPU")
if not top_k.is_cuda:
if torch.cuda.is_available():
top_k = top_k.cuda()
else:
raise RuntimeError("CUDA is not available but Triton kernel requires GPU")
# Get dimensions
batch_size, vocab_size = probs.shape
# Validate dimensions
assert vocab_size == 128256, f"Expected vocab_size=128256, got {vocab_size}"
assert top_k.shape == (batch_size,), f"Expected top_k shape ({batch_size},), got {top_k.shape}"
# Convert to required dtypes
probs = probs.to(torch.float32).contiguous()
top_k = top_k.to(torch.int32).contiguous()
# Allocate output
samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
# Generate random seed
seed = torch.randint(0, 2**31 - 1, (1,), device=probs.device).item()
# Configure grid
grid = (batch_size,)
# Determine block size based on vocab size
BLOCK_SIZE = 512 # Good for B200 GPU
# Launch optimized kernel
try:
optimized_top_k_kernel[grid](
probs_ptr=probs,
top_k_ptr=top_k,
samples_ptr=samples,
seed=seed,
batch_size=batch_size,
vocab_size=vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
)
except Exception as e:
# Fallback to simpler kernel if optimization fails
print(f"Warning: Optimized kernel failed with {e}, using fallback")
fallback_top_k_kernel[grid](
probs_ptr=probs,
top_k_ptr=top_k,
samples_ptr=samples,
seed=seed,
batch_size=batch_size,
vocab_size=vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
)
# Move result back to original device if needed
if original_probs_device != samples.device:
samples = samples.to(original_probs_device)
return samplesscrolls · 278 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON