gpt-5-2025-08-07 / triton4a1fba
gpt-5-2025-08-07_triton_4a1fba · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 245 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-4a1fba?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:6c942a0ebe33a55c2ebca69541d8c7b4a613104025e1fff7e14855a2e2f69148
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 = 8stages = 2
num_stages = 2Kernel source
main.py245 lines
import math
import torch
import triton
import triton.language as tl
VOCAB_SIZE = 151936
@triton.jit
def _top_p_sample_kernel(
probs_ptr, # *const float32 [B, V]
work_ptr, # *const float32 [B, V] - filtered probs (or same as probs)
top_p_ptr, # *const float32 [B]
out_ptr, # *mut int64 [B]
B, # int32: batch size
stride_p0, stride_p1, # strides for probs in elements
stride_w0, stride_w1, # strides for work in elements
seed, # uint32 seed for RNG
V: tl.constexpr, # vocab size (compile-time)
BLOCK: tl.constexpr # block size along vocab
):
pid = tl.program_id(axis=0)
if pid >= B:
return
# Load per-row top_p threshold
p = tl.load(top_p_ptr + pid)
# Compute base pointers for this row
row_p_ptr = probs_ptr + pid * stride_p0
row_w_ptr = work_ptr + pid * stride_w0
# Constants
eps = 1e-9
big_pos = 1e20
# Branch: p <= 0 => argmax over original probabilities
if p <= 0.0:
best_val = tl.full((), -float("inf"), dtype=tl.float32) # scalar
best_idx = tl.full((), 0, dtype=tl.int32) # scalar
for start in tl.static_range(0, V, BLOCK):
offs = start + tl.arange(0, BLOCK)
mask = offs < V
vals = tl.load(row_p_ptr + offs * stride_p1, mask=mask, other=-float("inf"))
# block max and first index achieving it
block_max = tl.max(vals, axis=0) # scalar
eq = vals == block_max
idxs = offs
cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
block_idx = tl.min(cand, axis=0) # scalar
# update global best (prefer smaller index on ties)
better = block_max > best_val
equal = block_max == best_val
take_idx = tl.where(equal, block_idx < best_idx, better)
best_idx = tl.where(take_idx, block_idx, best_idx)
best_val = tl.where(take_idx, block_max, best_val)
tl.store(out_ptr + pid, best_idx.to(tl.int64))
return
# Branch: p > 0 => sample from 'work' distribution using A-Res
# If p >= 1, 'work' should be equal to 'probs'. If 0 < p < 1, 'work' is top-p filtered.
best_r = tl.full((), big_pos, dtype=tl.float32) # scalar
best_i = tl.full((), 0, dtype=tl.int32) # scalar
sum_w = tl.full((), 0.0, dtype=tl.float32) # scalar
for start in tl.static_range(0, V, BLOCK):
offs = start + tl.arange(0, BLOCK)
mask = offs < V
w = tl.load(row_w_ptr + offs * stride_w1, mask=mask, other=0.0)
sum_w += tl.sum(w, axis=0)
# RNG: unique per (row, token). Keep offsets in 32-bit domain.
rng_offsets = (pid * V + offs).to(tl.int32)
u = tl.rand(seed, rng_offsets)
u = tl.maximum(u, eps)
denom = tl.maximum(w, eps)
r = -tl.log(u) / denom
# for zero-weight elements, set r to big_pos so they never win
r = tl.where(w > 0, r, big_pos)
# block min of r and first index achieving it
block_min_r = tl.min(r, axis=0) # scalar
eq = r == block_min_r
idxs = offs
cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
block_min_idx = tl.min(cand, axis=0) # scalar
better = block_min_r < best_r
best_r = tl.where(better, block_min_r, best_r)
best_i = tl.where(better, block_min_idx, best_i)
# If all weights were zero (degenerate), fallback to argmax over original probs
has_weight = sum_w > 0.0
if has_weight:
tl.store(out_ptr + pid, best_i.to(tl.int64))
else:
best_val = tl.full((), -float("inf"), dtype=tl.float32) # scalar
best_idx = tl.full((), 0, dtype=tl.int32) # scalar
for start in tl.static_range(0, V, BLOCK):
offs = start + tl.arange(0, BLOCK)
mask = offs < V
vals = tl.load(row_p_ptr + offs * stride_p1, mask=mask, other=-float("inf"))
block_max = tl.max(vals, axis=0) # scalar
eq = vals == block_max
idxs = offs
cand = tl.where(eq, idxs, tl.full([BLOCK], V, dtype=tl.int32))
block_idx = tl.min(cand, axis=0) # scalar
better = block_max > best_val
equal = block_max == best_val
take_idx = tl.where(equal, block_idx < best_idx, better)
best_idx = tl.where(take_idx, block_idx, best_idx)
best_val = tl.where(take_idx, block_max, best_val)
tl.store(out_ptr + pid, best_idx.to(tl.int64))
def _as_cuda(t: torch.Tensor) -> torch.Tensor:
if t.is_cuda:
return t
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required but not available. Cannot move CPU tensors to GPU.")
return t.cuda(non_blocking=True)
def _check_inputs(probs: torch.Tensor, top_p: torch.Tensor):
if probs.dtype != torch.float32:
raise TypeError(f"probs must be float32, got {probs.dtype}")
if probs.dim() != 2:
raise ValueError(f"probs must be 2D [batch, vocab], got shape {tuple(probs.shape)}")
if probs.shape[1] != VOCAB_SIZE:
raise ValueError(f"Expected vocab_size={VOCAB_SIZE}, got {probs.shape[1]}")
if top_p.dtype not in (torch.float32, torch.float64):
raise TypeError(f"top_p must be float tensor, got {top_p.dtype}")
if top_p.dim() != 1:
raise ValueError(f"top_p must be 1D [batch], got shape {tuple(top_p.shape)}")
if top_p.shape[0] != probs.shape[0]:
raise ValueError(f"top_p length {top_p.shape[0]} must match batch_size {probs.shape[0]}")
def _build_top_p_filtered_work(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
"""
Build a filtered copy of probs for top-p sampling:
- For rows with 0 < p < 1: zero out tokens outside the nucleus (HuggingFace-style mask).
- For rows with p <= 0 or p >= 1: leave as original probs.
All operations are on the same device as probs (expected GPU).
"""
B, V = probs.shape
work = probs.clone()
mask_rows = (top_p > 0.0) & (top_p < 1.0)
if mask_rows.any():
sel_probs = probs[mask_rows] # [B_sel, V]
sel_p = top_p[mask_rows].unsqueeze(1) # [B_sel, 1]
# Sort descending
vals, idxs = torch.sort(sel_probs, dim=1, descending=True) # [B_sel, V], [B_sel, V]
cdf = torch.cumsum(vals, dim=1) # [B_sel, V]
# HF-style nucleus mask: remove tokens strictly after the first that crosses p
to_remove = cdf > sel_p
to_remove_shifted = torch.zeros_like(to_remove)
to_remove_shifted[:, 1:] = to_remove[:, :-1]
keep_sorted = ~to_remove_shifted
filtered_sorted = vals * keep_sorted.to(vals.dtype)
# Scatter back into original index space
work_sel = torch.zeros_like(sel_probs)
work_sel.scatter_(1, idxs, filtered_sorted)
work[mask_rows] = work_sel
return work
@torch.no_grad()
def run(*args, **kwargs):
"""
Entry point: top_p_sampling_from_probs_v151936
Inputs:
- probs: [batch, 151936] float32 probabilities (after softmax)
- top_p: [batch] float32 cumulative probability thresholds
Output:
- samples: [batch] int64 sampled token indices
"""
# Extract args
if len(args) >= 2:
probs, top_p = args[0], args[1]
else:
probs = kwargs.get("probs", None)
top_p = kwargs.get("top_p", None)
if probs is None or top_p is None:
raise ValueError("run expects arguments (probs, top_p) either as positional or keyword.")
_check_inputs(probs, top_p)
# Preserve original device of probs
orig_device_probs = probs.device
# Ensure CUDA tensors
probs_cuda = _as_cuda(probs.contiguous())
top_p_cuda = _as_cuda(top_p.to(torch.float32).contiguous())
B, V = probs_cuda.shape
assert V == VOCAB_SIZE, f"Expected vocab={VOCAB_SIZE}, got {V}"
# Build filtered work tensor for 0 < p < 1 rows (GPU)
work = _build_top_p_filtered_work(probs_cuda, top_p_cuda)
# Output buffer on GPU
out_cuda = torch.empty((B,), dtype=torch.int64, device=probs_cuda.device)
# Strides in elements (not bytes)
stride_p0, stride_p1 = probs_cuda.stride()
stride_w0, stride_w1 = work.stride()
# Random seed: uint32
seed = torch.randint(0, 2**31 - 1, (1,), device=probs_cuda.device, dtype=torch.int64).item()
seed = int(seed & 0xFFFFFFFF)
# Kernel launch configuration - tuned for large V on B200
BLOCK = 4096 # tile over vocab
num_warps = 8
num_stages = 2
grid = lambda META: (B,)
_top_p_sample_kernel[grid](
probs_cuda,
work,
top_p_cuda,
out_cuda,
B,
stride_p0, stride_p1,
stride_w0, stride_w1,
seed,
V=VOCAB_SIZE,
BLOCK=BLOCK,
num_warps=num_warps,
num_stages=num_stages,
)
# Move result back to original device of probs
if orig_device_probs.type == "cpu":
return out_cuda.cpu()
return out_cudascrolls · 245 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON