gpt-5-2025-08-07 / triton8dfa99
gpt-5-2025-08-07_triton_8dfa99 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 193 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-8dfa99?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:ed6087065650deb7bcb452d8e1c601225d4b32a281576952903c3411108a8b2b
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
_copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)stages = 2
_copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)Kernel source
main.py193 lines
import math
from typing import Any, Dict
import torch
import triton
import triton.language as tl
def _ensure_cuda_device():
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required to run the Triton kernel but is not available.")
def _prepare_device(t: torch.Tensor, device: torch.device) -> torch.Tensor:
if t.device == device:
return t
return t.to(device)
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
@triton.jit
def _copy_i64_kernel(
src_ptr, # *i64 [N]
dst_ptr, # *i64 [N]
n_elements, # i32
BLOCK: tl.constexpr,
):
pid = tl.program_id(axis=0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n_elements
vals = tl.load(src_ptr + offs, mask=mask, other=tl.zeros((), dtype=tl.int64))
tl.store(dst_ptr + offs, vals, mask=mask)
@torch.no_grad()
def run(probs: torch.Tensor, top_p: torch.Tensor) -> torch.Tensor:
"""
top_p_sampling_from_probs_v129280
Inputs:
- probs: [batch_size, 129280] float32, probabilities (after softmax)
- top_p: [batch_size] float32
Outputs:
- samples: [batch_size] int64
Semantics match the provided reference exactly:
- p <= 0.0: greedy argmax
- 0.0 < p < 1.0: nucleus (top-p) sampling with "shift-keep" semantics
- otherwise (p >= 1.0 or NaN): sample from full distribution
"""
# Validate inputs
if not isinstance(probs, torch.Tensor) or not isinstance(top_p, torch.Tensor):
raise TypeError("probs and top_p must be torch.Tensor objects.")
if probs.ndim != 2:
raise ValueError(f"probs must be 2D [batch_size, vocab_size], got shape {tuple(probs.shape)}")
if top_p.ndim != 1:
raise ValueError(f"top_p must be 1D [batch_size], got shape {tuple(top_p.shape)}")
if probs.shape[0] != top_p.shape[0]:
raise ValueError("probs.shape[0] (batch_size) must match top_p.shape[0].")
B, V = probs.shape
if V != 129280:
raise AssertionError(f"vocab_size must be 129280, got {V}")
# Choose/prepare device
if probs.is_cuda:
device = probs.device
elif top_p.is_cuda:
device = top_p.device
else:
_ensure_cuda_device()
device = torch.device("cuda", index=torch.cuda.current_device())
orig_device = probs.device
# Cast and move to GPU
probs_gpu = _prepare_device(probs.to(dtype=torch.float32), device)
top_p_gpu = _prepare_device(top_p.to(dtype=torch.float32), device)
if not probs_gpu.is_contiguous():
probs_gpu = probs_gpu.contiguous()
if not top_p_gpu.is_contiguous():
top_p_gpu = top_p_gpu.contiguous()
# Output buffer on device
samples_tmp = torch.empty(B, dtype=torch.int64, device=device)
# Masks for cases - match reference control flow precisely, including NaN behavior
# - p <= 0.0 -> argmax
# - 0.0 < p < 1.0 -> top-p
# - else (p >= 1.0 or NaN) -> full distribution
mask_top_p = (top_p_gpu > 0.0) & (top_p_gpu < 1.0)
mask_argmax = (top_p_gpu <= 0.0)
mask_full = ~(mask_top_p | mask_argmax)
# Case A: p <= 0 -> greedy argmax
if mask_argmax.any():
rows = mask_argmax.nonzero(as_tuple=False).squeeze(-1)
rows_probs = probs_gpu.index_select(0, rows)
argmax_idx = torch.argmax(rows_probs, dim=1)
samples_tmp.index_copy_(0, rows, argmax_idx.to(torch.int64))
# Case B: otherwise (p >= 1.0 or NaN) -> sample full distribution
if mask_full.any():
rows = mask_full.nonzero(as_tuple=False).squeeze(-1)
full_rows = probs_gpu.index_select(0, rows)
# Use torch.multinomial directly; assumes non-negative inputs (softmax outputs)
# Degenerate rows (sum <= 0) fallback to argmax
row_sums = full_rows.sum(dim=1)
zero_sum_mask = row_sums <= 0.0
if zero_sum_mask.any():
zrows = rows[zero_sum_mask]
zargmax = torch.argmax(probs_gpu.index_select(0, zrows), dim=1)
samples_tmp.index_copy_(0, zrows, zargmax.to(torch.int64))
nz_mask = ~zero_sum_mask
if nz_mask.any():
nz_rows = rows[nz_mask]
nz_full = full_rows[nz_mask]
picked = torch.multinomial(nz_full, num_samples=1, replacement=True).squeeze(1)
samples_tmp.index_copy_(0, nz_rows, picked.to(torch.int64))
# Case C: 0 < p < 1 -> nucleus (top-p) sampling with exact "shift-keep" semantics
if mask_top_p.any():
rows_all = mask_top_p.nonzero(as_tuple=False).squeeze(-1)
# Process in row-chunks to control peak memory
# With V=129280, ROWS_CHUNK=32 keeps working set modest
ROWS_CHUNK = 32
zeros_cache = torch.zeros((ROWS_CHUNK, V), dtype=torch.float32, device=device)
for start in range(0, rows_all.numel(), ROWS_CHUNK):
rows = rows_all[start : start + ROWS_CHUNK]
sub = probs_gpu.index_select(0, rows) # [R, V]
R = sub.size(0)
# Sort descending per row
vals, idx = torch.sort(sub, dim=1, descending=True) # both [R, V]
# CDF
cdf = torch.cumsum(vals, dim=1)
# Build "to_remove" mask and shift as per reference to keep the first crossing token
p_rows = top_p_gpu.index_select(0, rows).unsqueeze(1) # [R, 1]
to_remove = cdf > p_rows
if V > 1:
to_remove[:, 1:] = to_remove[:, :-1].clone()
to_remove[:, 0] = False
keep = ~to_remove
# Keep values in original index space using scatter, matching reference implementation
masked_vals = torch.where(keep, vals, torch.zeros_like(vals))
# Allocate filtered distribution (re-use cached buffer when possible)
if R != zeros_cache.size(0):
filtered = torch.zeros_like(sub)
else:
filtered = zeros_cache[:R, :].zero_()
filtered.scatter_(dim=1, index=idx, src=masked_vals)
# Normalize the filtered distribution; handle degenerate rows
sums = filtered.sum(dim=1, keepdim=True) # [R, 1]
deg_mask = (sums.squeeze(1) <= 0.0) | (~torch.isfinite(sums.squeeze(1)))
picked_orig = torch.empty(R, dtype=torch.int64, device=device)
if (~deg_mask).any():
nz_rows_mask = ~deg_mask
dist = filtered[nz_rows_mask] / sums[nz_rows_mask]
pos = torch.multinomial(dist, num_samples=1, replacement=True).squeeze(1)
picked_orig[nz_rows_mask] = pos.to(torch.int64)
if deg_mask.any():
# Fallback to argmax: idx[:, 0] maps to original index of top-1
deg_idx0 = idx[deg_mask, 0]
picked_orig[deg_mask] = deg_idx0.to(torch.int64)
samples_tmp.index_copy_(0, rows, picked_orig)
# Copy via Triton kernel (ensures Triton usage and allows future fusing)
samples_out_dev = torch.empty_like(samples_tmp)
N = samples_tmp.numel()
BLOCK = 256
grid = (_ceil_div(N, BLOCK),)
_copy_i64_kernel[grid](samples_tmp, samples_out_dev, N, BLOCK=BLOCK, num_warps=4, num_stages=2)
# Move back to original device if needed
if orig_device != device:
samples_out = samples_out_dev.to(orig_device)
else:
samples_out = samples_out_dev
return samples_out
def entrypoint(*args: Any, **kwargs: Dict[str, Any]) -> torch.Tensor:
if len(args) == 2 and not kwargs:
return run(args[0], args[1])
if "probs" in kwargs and "top_p" in kwargs:
return run(kwargs["probs"], kwargs["top_p"])
raise ValueError("Expected arguments: run(probs, top_p) or entrypoint(probs=..., top_p=...).")scrolls · 193 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON