gpt-5-2025-08-07 / triton44f7ae
gpt-5-2025-08-07_triton_44f7ae · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 243 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-44f7ae?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:3c11b901ec5b15f8a93b72673b6125fdf63413ff7b4b055983827bb4e6369784
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 = 8
num_warps=8,stages = 3
num_stages=3,Kernel source
main.py243 lines
import math
import torch
import triton
import triton.language as tl
VOCAB_SIZE = 128256
@triton.jit
def _top_p_sample_sorted_kernel(
vals_ptr, # float32 [B, V] sorted descending per row
idx_ptr, # int64 [B, V] corresponding original indices
top_p_ptr, # float32 [B]
rand_ptr, # float32 [B] uniform in [0, 1)
out_ptr, # int64 [B]
batch_size: tl.constexpr,
vocab_size: tl.constexpr,
CHUNK_SIZE: tl.constexpr,
CHUNKS: tl.constexpr,
):
pid = tl.program_id(axis=0)
if pid >= batch_size:
return
# Base pointers for this row
row_vals_ptr = vals_ptr + pid * vocab_size
row_idx_ptr = idx_ptr + pid * vocab_size
# Per-row parameters
p = tl.load(top_p_ptr + pid)
u01 = tl.load(rand_ptr + pid)
# Degenerate case: p <= 0.0 -> argmax (sorted -> first element)
if p <= 0.0:
best_idx = tl.load(row_idx_ptr + 0)
tl.store(out_ptr + pid, best_idx)
return
# Compile-time arange for chunk indexing
ar = tl.arange(0, CHUNK_SIZE)
# First pass: find truncation boundary (if 0 < p < 1), and total mass
total_mass = tl.full((), 0.0, dtype=tl.float32)
cum_before = tl.full((), 0.0, dtype=tl.float32)
found = tl.full((), False, dtype=tl.int1)
bound_chunk = tl.full((), 0, dtype=tl.int32) # chunk id where we cross p
bound_i_local = tl.full((), 0, dtype=tl.int32) # local index within bound_chunk
t_mass = tl.full((), 0.0, dtype=tl.float32) # truncated mass up to boundary (inclusive)
big_i_vec = tl.full([CHUNK_SIZE], 2147483647, dtype=tl.int32) # for reductions
for j in range(CHUNKS):
base = j * CHUNK_SIZE
offs = base + ar
valid = offs < vocab_size
v = tl.load(row_vals_ptr + offs, mask=valid, other=0.0)
s_chunk = tl.sum(v, axis=0)
# Always accumulate total mass (for treat_full case)
total_mass = total_mass + s_chunk
# Check if we need to search the crossing in this chunk
truncated = p < 1.0
need = truncated & (~found)
# Compute prefix sums within this chunk (masked by valid) plus cumulative before
pref = tl.cumsum(v, axis=0) + cum_before
# Determine if crossing happens within this chunk
is_cross = pref > p
# Replace tl.any with reduction to float and comparison
any_cross = tl.max(tl.where(is_cross, 1.0, 0.0), axis=0) > 0.0
any_cross = need & any_cross
# First crossing index within this chunk (if any)
idx_first = tl.min(tl.where(is_cross, ar, big_i_vec), axis=0)
# Mass at crossing
pref_selected = tl.sum(tl.where(ar == idx_first, pref, 0.0), axis=0)
# Update boundary if we found crossing here
bound_chunk = tl.where(any_cross, tl.full((), j, dtype=tl.int32), bound_chunk)
bound_i_local = tl.where(any_cross, idx_first, bound_i_local)
t_mass = tl.where(any_cross, pref_selected, t_mass)
found = found | any_cross
# If still not found and truncating, accumulate chunk mass into cum_before
cum_before = tl.where(need & (~any_cross), cum_before + s_chunk, cum_before)
# If not truncating or we never crossed p, treat as full distribution
treat_full = (~(p < 1.0)) | (~found)
t_mass = tl.where(treat_full, total_mass, t_mass)
# If truncated mass is non-positive (degenerate), fall back to argmax
if t_mass <= 0.0:
best_idx = tl.load(row_idx_ptr + 0)
tl.store(out_ptr + pid, best_idx)
return
# Sample u in [0, t_mass)
u = u01 * t_mass
# Second pass: locate sampled token by scanning allowed mass
acc = tl.full((), 0.0, dtype=tl.float32)
picked = tl.full((), False, dtype=tl.int1)
sel_off = tl.full((), 0, dtype=tl.int32)
for j in range(CHUNKS):
base = j * CHUNK_SIZE
offs = base + ar
valid = offs < vocab_size
v = tl.load(row_vals_ptr + offs, mask=valid, other=0.0)
# Build allow mask depending on truncation:
# allow_trunc = valid & ((j < bound_chunk) | ((j == bound_chunk) & (ar <= bound_i_local)))
j_scalar = tl.full((), j, dtype=tl.int32)
before = j_scalar < bound_chunk
at = j_scalar == bound_chunk
le_local = ar <= bound_i_local
allow_trunc = valid & (before | (at & le_local))
allow = tl.where(treat_full, valid, allow_trunc)
w = tl.where(allow, v, 0.0)
s_allow = tl.sum(w, axis=0)
pref = tl.cumsum(w, axis=0) + acc
cross = pref > u
any_cross = tl.max(tl.where(cross, 1.0, 0.0), axis=0) > 0.0
# Only consider first time we find the crossing
want_pick = (~picked) & any_cross
idx_first = tl.min(tl.where(cross, ar, big_i_vec), axis=0)
pos_global = base + idx_first
sel_off = tl.where(want_pick, pos_global, sel_off)
picked = picked | want_pick
# Update accumulator only if we still haven't picked
acc = tl.where(~picked, acc + s_allow, acc)
# If nothing picked due to numerical issues, default to first token
sel_off = tl.where(picked, sel_off, tl.full((), 0, dtype=tl.int32))
# Fetch original token index and store
tok_idx = tl.load(row_idx_ptr + sel_off)
tl.store(out_ptr + pid, tok_idx)
def _ensure_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
if t.device == device:
return t
if device.type == "cuda":
return t.to(device, non_blocking=True)
return t.cpu()
def _validate_and_prepare_inputs(probs: torch.Tensor, top_p: torch.Tensor):
if probs.dim() != 2:
raise ValueError(f"probs must be 2D [batch_size, vocab_size], got {tuple(probs.shape)}")
if probs.shape[1] != VOCAB_SIZE:
raise ValueError(f"vocab_size must be {VOCAB_SIZE}, got {probs.shape[1]}")
if top_p.dim() != 1:
raise ValueError(f"top_p must be 1D [batch_size], got {tuple(top_p.shape)}")
if top_p.shape[0] != probs.shape[0]:
raise ValueError(f"top_p batch size {top_p.shape[0]} does not match probs batch size {probs.shape[0]}")
if probs.dtype != torch.float32:
probs = probs.to(torch.float32)
if top_p.dtype != torch.float32:
top_p = top_p.to(torch.float32)
return probs, top_p
def _top_p_sample_impl(probs: torch.Tensor, top_p: torch.Tensor, generator: torch.Generator = None) -> torch.Tensor:
# Validate and cast types
probs, top_p = _validate_and_prepare_inputs(probs, top_p)
batch_size, vocab_size = probs.shape
# Device management
if probs.is_cuda or top_p.is_cuda:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU tensors were provided.")
# Ensure both inputs on the same CUDA device
target_device = probs.device if probs.is_cuda else top_p.device
if probs.is_cuda and top_p.is_cuda and probs.device != top_p.device:
target_device = probs.device
else:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run the Triton kernel, but no CUDA device is available.")
target_device = torch.device("cuda")
probs_gpu = _ensure_device(probs.contiguous(), target_device)
top_p_gpu = _ensure_device(top_p.contiguous(), target_device)
# Sort probabilities descending per row; get sorted values and original indices
sorted_vals, sorted_idx = torch.sort(probs_gpu, dim=1, descending=True, stable=False)
# Random uniforms per row
if generator is None:
rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device)
else:
try:
rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device, generator=generator)
except Exception:
rand = torch.rand((batch_size,), dtype=torch.float32, device=target_device)
# Output buffer
out = torch.empty((batch_size,), dtype=torch.int64, device=target_device)
# Launch Triton kernel: one program per row
CHUNK_SIZE = 1024 # tile size for vectorized scanning; good fit for B200
CHUNKS = (VOCAB_SIZE + CHUNK_SIZE - 1) // CHUNK_SIZE
grid = (batch_size,)
_top_p_sample_sorted_kernel[grid](
sorted_vals,
sorted_idx,
top_p_gpu,
rand,
out,
batch_size=batch_size,
vocab_size=vocab_size,
CHUNK_SIZE=CHUNK_SIZE,
CHUNKS=CHUNKS,
num_warps=8,
num_stages=3,
)
# Move result back to the original device of probs
out_final = out.to(probs.device) if probs.device.type != "cuda" else out
return out_final
def run(*args, **kwargs):
"""
Entry point. Usage:
samples = run(probs, top_p)
"""
if len(args) < 2 and not ("probs" in kwargs and "top_p" in kwargs):
raise ValueError("run requires 'probs' and 'top_p' arguments.")
probs = args[0] if len(args) > 0 else kwargs["probs"]
top_p = args[1] if len(args) > 1 else kwargs["top_p"]
generator = kwargs.get("generator", None)
return _top_p_sample_impl(probs, top_p, generator=generator)scrolls · 243 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON