gpt-5-2025-08-07 / triton7230f5
gpt-5-2025-08-07_triton_7230f5 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 229 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-7230f5?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:bdf3112b281d0909c51c10df0c5a799e42686adb033a18606b87af3c1f4561ab
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps = 4stages = 2
num_stages = 2Kernel source
main.py229 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _sample_from_topk_kernel(
topk_vals_ptr, # float32* [N_valid, Kmax]
topk_idx_ptr, # int32* [N_valid, Kmax]
sizes_ptr, # int32* [N_valid]
rand_ptr, # float32* [N_valid], uniform in [0, sum(topk_vals_i))
row_map_ptr, # int32* [N_valid], maps local row -> original batch row
out_ptr, # int64* [batch_size], output indices
Kmax: tl.constexpr, # padded max-k across valid rows
BLOCK: tl.constexpr, # tile size along K dimension, e.g., 256
):
pid = tl.program_id(0)
# Per-row metadata (0-D scalars)
k_size = tl.load(sizes_ptr + pid)
r = tl.load(rand_ptr + pid)
row_out = tl.load(row_map_ptr + pid)
# Running state (0-D scalars)
acc = tl.zeros((), dtype=tl.float32) # accumulated sum before current block
found = tl.zeros((), dtype=tl.int32) # 0/1 flag
found_idx_global = tl.full((), -1, dtype=tl.int32)
row_base = pid * Kmax
arange = tl.arange(0, BLOCK)
# Iterate blocks across K dimension with compile-time unrolling
for start in tl.static_range(0, Kmax, BLOCK):
offs = start + arange
# Valid elements within this block for this row
valid = offs < k_size
vals = tl.load(topk_vals_ptr + row_base + offs, mask=valid, other=0.0)
# Sum of this block
block_sum = tl.sum(vals, axis=0)
# Will the crossing happen within this block?
cross_in_block = (found == 0) & (acc + block_sum >= r)
# Sequential search within the block if needed. Avoid tensor indexing by scalar;
# instead do masked scalar loads directly from memory.
rem = k_size - start
sel = tl.full((), -1, tl.int32)
run_sum = acc
for j in tl.static_range(0, BLOCK):
j_mask = cross_in_block & (j < rem) & (sel < 0)
v = tl.load(topk_vals_ptr + row_base + (start + j), mask=j_mask, other=0.0)
run_sum = tl.where(j_mask, run_sum + v, run_sum)
crossed = j_mask & (run_sum >= r) & (sel < 0)
sel = tl.where(crossed, tl.full((), j, tl.int32), sel)
block_found = sel >= 0
found = tl.where(block_found & (found == 0), tl.full((), 1, tl.int32), found)
found_idx_global = tl.where(
block_found & (found_idx_global < 0),
tl.full((), start, tl.int32) + sel,
found_idx_global,
)
# If still not found, add this block's sum to acc
acc = tl.where(found == 0, acc + block_sum, acc)
# Fallback to last valid index if numerical corner-case prevented finding a crossing
last_idx = tl.where(k_size > 0, k_size - tl.full((), 1, tl.int32), tl.full((), 0, tl.int32))
final_pos = tl.where(found_idx_global < 0, last_idx, found_idx_global)
# Gather original token index and store
tok_i32 = tl.load(topk_idx_ptr + row_base + final_pos)
tok_i64 = tok_i32.to(tl.int64)
tl.store(out_ptr + row_out, tok_i64)
def _ensure_cuda_tensor(t: torch.Tensor, device: torch.device):
if t.device.type == "cuda":
if t.device != device:
return t.to(device)
return t
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required but not available. Cannot move CPU tensors to GPU.")
return t.to(device)
@torch.no_grad()
def run(probs, top_k):
"""
Triton-accelerated top-k sampling from probability distributions.
Args:
probs: [batch_size, 128256] float32, probabilities after softmax.
top_k: [batch_size] int32, per-row top-k values. If 0 < k < vocab_size, restrict to top-k tokens,
renormalize implicitly via weighted sampling and sample. Otherwise sample from the full distribution.
Returns:
samples: [batch_size] int64, sampled token indices.
"""
# Handle both args and kwargs robustly
if isinstance(probs, dict):
probs = probs.get("probs", None)
if isinstance(top_k, dict):
top_k = top_k.get("top_k", None)
if probs is None or top_k is None:
raise ValueError("Both 'probs' and 'top_k' must be provided.")
# Basic validation and types
if probs.ndim != 2:
raise ValueError(f"probs must be 2D [batch_size, vocab_size], got shape {tuple(probs.shape)}")
if top_k.ndim != 1:
raise ValueError(f"top_k must be 1D [batch_size], got shape {tuple(top_k.shape)}")
batch_size, vocab_size = probs.shape
if vocab_size != 128256:
raise AssertionError(f"vocab_size must be 128256, got {vocab_size}")
# Convert dtypes exactly as in the reference
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
# Device management
orig_device = probs.device
if torch.cuda.is_available():
if probs.is_cuda:
work_device = probs.device
elif top_k.is_cuda:
work_device = top_k.device
else:
work_device = torch.device("cuda")
else:
if probs.is_cuda or top_k.is_cuda:
raise RuntimeError("CUDA is not available but a tensor is on GPU.")
raise RuntimeError("CUDA is required for this Triton kernel but is not available.")
# Move to working CUDA device if needed
probs_gpu = _ensure_cuda_tensor(probs, work_device)
top_k_gpu = _ensure_cuda_tensor(top_k, work_device)
# Output buffer on GPU
samples_gpu = torch.empty((batch_size,), dtype=torch.int64, device=work_device)
# Mask rows by validity of k
V = vocab_size
k_valid_mask = (top_k_gpu > 0) & (top_k_gpu < V)
valid_rows = torch.nonzero(k_valid_mask, as_tuple=False).flatten()
invalid_rows = torch.nonzero(~k_valid_mask, as_tuple=False).flatten()
# Handle invalid-k rows: sample from full distribution using torch.multinomial (GPU)
if invalid_rows.numel() > 0:
probs_invalid = probs_gpu.index_select(0, invalid_rows).contiguous()
sampled_invalid = torch.multinomial(probs_invalid, num_samples=1, replacement=True).squeeze(1).to(torch.int64)
samples_gpu.index_copy_(0, invalid_rows, sampled_invalid)
# Handle valid-k rows with Triton kernel
if valid_rows.numel() > 0:
# Gather valid rows
probs_valid = probs_gpu.index_select(0, valid_rows).contiguous()
k_vals = top_k_gpu.index_select(0, valid_rows) # [N_valid] int32
Kmax = int(k_vals.max().item())
N_valid = probs_valid.size(0)
# Compute top-Kmax once for all valid rows (sorted desc)
topk_vals_padded, topk_idx_padded = torch.topk(probs_valid, Kmax, dim=1, largest=True, sorted=True)
topk_vals_padded = topk_vals_padded.contiguous()
topk_idx_padded = topk_idx_padded.to(torch.int32).contiguous() # Triton expects int32
# Compute per-row sums across the first k_i entries only
ar = torch.arange(Kmax, device=work_device, dtype=torch.int32).unsqueeze(0) # [1, Kmax]
sizes_broadcast = k_vals.unsqueeze(1) # [N_valid, 1]
mask2d = (ar < sizes_broadcast) # [N_valid, Kmax], bool
sums = (topk_vals_padded * mask2d.to(topk_vals_padded.dtype)).sum(dim=1) # [N_valid]
# Safety: if any sum is 0 (shouldn't happen), fall back to full distribution for those rows
zero_sum_mask = sums <= 0
if torch.any(zero_sum_mask):
fix_rows_local = torch.nonzero(zero_sum_mask, as_tuple=False).flatten()
if fix_rows_local.numel() > 0:
fix_rows_global = valid_rows.index_select(0, fix_rows_local)
probs_fix = probs_gpu.index_select(0, fix_rows_global).contiguous()
sampled_fix = torch.multinomial(probs_fix, num_samples=1, replacement=True).squeeze(1).to(torch.int64)
samples_gpu.index_copy_(0, fix_rows_global, sampled_fix)
keep_mask = ~zero_sum_mask
if torch.any(keep_mask):
keep_idx = torch.nonzero(keep_mask, as_tuple=False).flatten()
topk_vals_padded = topk_vals_padded.index_select(0, keep_idx).contiguous()
topk_idx_padded = topk_idx_padded.index_select(0, keep_idx).contiguous()
k_vals = k_vals.index_select(0, keep_idx).contiguous()
valid_rows_kernel = valid_rows.index_select(0, keep_idx).contiguous()
sums = sums.index_select(0, keep_idx).contiguous()
N_valid_kernel = valid_rows_kernel.numel()
else:
N_valid_kernel = 0
else:
valid_rows_kernel = valid_rows
N_valid_kernel = N_valid
if N_valid_kernel > 0:
# Prepare random thresholds in [0, sums)
rands = torch.rand((N_valid_kernel,), dtype=torch.float32, device=work_device) * sums
# Launch Triton kernel
grid = (N_valid_kernel,)
# Tuned params for B200
num_warps = 4
num_stages = 2
BLOCK = 256 # good trade-off for memory coalescing vs. register pressure
# Row mapping back to global batch indices
row_map = valid_rows_kernel.to(torch.int32).contiguous()
_sample_from_topk_kernel[grid](
topk_vals_padded,
topk_idx_padded,
k_vals,
rands,
row_map,
samples_gpu,
Kmax=Kmax,
BLOCK=BLOCK,
num_warps=num_warps,
num_stages=num_stages,
)
# Move result back to original device if needed
samples = samples_gpu if orig_device.type == "cuda" else samples_gpu.to(orig_device)
return samplesscrolls · 229 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON