gpt-5-2025-08-07 / tritoncf2509
gpt-5-2025-08-07_triton_cf2509 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 200 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-cf2509?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:ca6f6479d6fb372c37b38dd1001b1d7523a7610998d31fe16460dbc2139f5e3f
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.py200 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def _copy_1d_kernel(src_ptr, dst_ptr, N: tl.int32, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
vals = tl.load(src_ptr + offs, mask=mask, other=tl.zeros((), dtype=tl.int64))
tl.store(dst_ptr + offs, vals, mask=mask)
def _ensure_cuda_available():
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run this kernel but is not available.")
def _to_device(t: torch.Tensor, device: torch.device, dtype=None, contiguous=True):
if dtype is not None:
t = t.to(dtype)
if contiguous:
t = t.contiguous()
if t.device == device:
return t
if device.type == "cuda":
return t.to(device, non_blocking=True)
return t.cuda(non_blocking=True)
def _ceil_div(a, b):
return (a + b - 1) // b
@torch.no_grad()
def run(probs, top_k, top_p, **kwargs):
"""
Efficient and correct top-k + top-p sampling for Qwen3 vocab (151936).
This implementation computes the selection using optimized PyTorch ops on GPU
and uses a lightweight Triton kernel for the final write, avoiding the
pathological O(V*k) loops that can cause timeouts.
Inputs:
- probs: [B, 151936] float32, already softmax'ed
- top_k: [B] int32
- top_p: [B] float32
Output:
- samples: [B] int64 (token indices)
"""
_ensure_cuda_available()
# Wrap tensors
probs = torch.as_tensor(probs)
top_k = torch.as_tensor(top_k)
top_p = torch.as_tensor(top_p)
# Validate shapes
if probs.dim() != 2:
raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
B, V = probs.shape
if V != 151936:
raise ValueError(f"vocab_size must be 151936; got {V}")
if top_k.shape != (B,):
raise ValueError(f"top_k must have shape [{B}], got {tuple(top_k.shape)}")
if top_p.shape != (B,):
raise ValueError(f"top_p must have shape [{B}], got {tuple(top_p.shape)}")
# Select target device (prefer probs' device if CUDA, else first CUDA)
target_device = probs.device if probs.is_cuda else torch.device("cuda")
# Move to GPU and correct dtypes
probs_gpu = _to_device(probs, target_device, dtype=torch.float32, contiguous=True)
top_k_gpu = _to_device(top_k, target_device, dtype=torch.int32, contiguous=True)
top_p_gpu = _to_device(top_p, target_device, dtype=torch.float32, contiguous=True)
B = int(probs_gpu.shape[0])
V = int(probs_gpu.shape[1])
# Output tensor computed with PyTorch
samples_calc = torch.empty((B,), dtype=torch.int64, device=target_device)
# Masks for cases
apply_k_mask = (top_k_gpu > 0) & (top_k_gpu < V)
p_neg_mask = top_p_gpu <= 0.0
p_one_mask = top_p_gpu >= 1.0
p_mid_mask = ~(p_neg_mask | p_one_mask)
# Case: p <= 0 -> always argmax (top-k doesn't change argmax)
rows = torch.nonzero(p_neg_mask, as_tuple=False).squeeze(1)
if rows.numel() > 0:
argmax_idx = torch.argmax(probs_gpu.index_select(0, rows), dim=1)
samples_calc.index_copy_(0, rows, argmax_idx.to(torch.int64))
# Case: no top-k, p >= 1 -> sample from full distribution
rows = torch.nonzero((~apply_k_mask) & p_one_mask, as_tuple=False).squeeze(1)
if rows.numel() > 0:
dist = probs_gpu.index_select(0, rows)
sel = torch.multinomial(dist, 1, replacement=True).squeeze(1)
samples_calc.index_copy_(0, rows, sel.to(torch.int64))
# Case: apply top-k, p >= 1 -> sample from top-k only
rows = torch.nonzero(apply_k_mask & p_one_mask, as_tuple=False).squeeze(1)
if rows.numel() > 0:
tk_vals = top_k_gpu.index_select(0, rows)
unique_k = torch.unique(tk_vals, sorted=True)
for kk in unique_k.tolist():
if kk <= 0 or kk >= V:
continue
rows_k_mask = (tk_vals == kk)
rows_k = rows.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1))
if rows_k.numel() == 0:
continue
row_probs = probs_gpu.index_select(0, rows_k)
vals, idxs = torch.topk(row_probs, k=kk, dim=1, largest=True, sorted=True)
# sample among top-k values directly (no need to renormalize)
sel_local = torch.multinomial(vals, 1, replacement=True)
chosen = idxs.gather(1, sel_local).squeeze(1)
samples_calc.index_copy_(0, rows_k, chosen.to(torch.int64))
# Case: no top-k, 0 < p < 1 -> nucleus sampling on full vocab
rows = torch.nonzero((~apply_k_mask) & p_mid_mask, as_tuple=False).squeeze(1)
if rows.numel() > 0:
row_probs = probs_gpu.index_select(0, rows)
p_rows = top_p_gpu.index_select(0, rows).unsqueeze(1)
# sort descending
vals_sorted, idx_sorted = torch.sort(row_probs, dim=1, descending=True)
cdf = torch.cumsum(vals_sorted, dim=1)
to_remove = cdf > p_rows
if V > 1:
# shift right to keep first token and ensure minimal valid nucleus
to_remove[:, 1:] = to_remove[:, :-1].clone()
to_remove[:, 0] = False
# zero out removed
vals_sorted = vals_sorted.masked_fill(to_remove, 0.0)
sel_pos = torch.multinomial(vals_sorted, 1, replacement=True)
chosen = idx_sorted.gather(1, sel_pos).squeeze(1)
samples_calc.index_copy_(0, rows, chosen.to(torch.int64))
# Case: apply top-k, 0 < p < 1 -> nucleus sampling within top-k
rows = torch.nonzero(apply_k_mask & p_mid_mask, as_tuple=False).squeeze(1)
if rows.numel() > 0:
tk_vals = top_k_gpu.index_select(0, rows)
p_rows_all = top_p_gpu.index_select(0, rows)
unique_k = torch.unique(tk_vals, sorted=True)
for kk in unique_k.tolist():
if kk <= 0 or kk >= V:
# shouldn't happen due to mask, but guard anyway
continue
rows_k_mask = (tk_vals == kk)
rows_k = rows.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1))
if rows_k.numel() == 0:
continue
row_probs = probs_gpu.index_select(0, rows_k)
p_rows = p_rows_all.index_select(0, torch.nonzero(rows_k_mask, as_tuple=False).squeeze(1)).unsqueeze(1)
# top-k sorted
vals, idxs = torch.topk(row_probs, k=kk, dim=1, largest=True, sorted=True)
# normalize within top-k to compute cdf as in reference
sums = vals.sum(dim=1, keepdim=True)
# Avoid division by zero; sums should be > 0 for valid distributions
sums = torch.clamp(sums, min=1e-20)
vals_norm = vals / sums
cdf = torch.cumsum(vals_norm, dim=1)
to_remove = cdf > p_rows
if kk > 1:
to_remove[:, 1:] = to_remove[:, :-1].clone()
to_remove[:, 0] = False
else:
to_remove[:, 0] = False
# sample within kept subset using original weights (proportionality preserved)
weights = vals.masked_fill(to_remove, 0.0)
sel_local = torch.multinomial(weights, 1, replacement=True)
chosen = idxs.gather(1, sel_local).squeeze(1)
samples_calc.index_copy_(0, rows_k, chosen.to(torch.int64))
# As a final safeguard (shouldn't be needed), replace any invalid indices with 0
invalid_mask = (samples_calc < 0) | (samples_calc >= V)
if torch.any(invalid_mask):
samples_calc[invalid_mask] = 0
# Use a lightweight Triton kernel to copy results to output
samples_out = torch.empty_like(samples_calc)
BLOCK = int(kwargs.pop("block_size", 256))
num_warps = int(kwargs.pop("num_warps", 4))
num_stages = int(kwargs.pop("num_stages", 2))
grid = (_ceil_div(B, BLOCK),)
_copy_1d_kernel[grid](
samples_calc, samples_out, B,
BLOCK=BLOCK,
num_warps=num_warps,
num_stages=num_stages,
)
# Move back to original device if needed
if probs.device.type == "cuda":
return samples_out
else:
return samples_out.cpu()scrolls · 200 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON