flashinfer / wrappere53f28
flashinfer_wrapper_e53f28 · flashinfer · python · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 53 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-e53f28?include=source"interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA A100, NVIDIA B200, NVIDIA H100, NVIDIA H20, NVIDIA H200
architecturesunknown
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:099b6b2717fb668f591e4dcee6260e0156476163ec49fa8916a5e5e067294182
license declaredApache-2.0
license concludedApache-2.0
authorsflashinfer
imported2026-08-16
Kernel source
main.py53 lines
import torch
import flashinfer
@torch.no_grad()
def run(probs, top_k, top_p):
batch_size, vocab_size = probs.shape
device = probs.device
assert vocab_size == 202048
probs = probs.to(torch.float32)
# Use the largest k across the batch (all workloads use same k=1000)
k = int(top_k.max().item())
k = max(1, min(k, vocab_size))
# Get the top-k tokens and their original (non-renormed) probabilities.
# torch.topk returns values sorted descending, so topk_vals[i] is already
# sorted from highest to lowest probability for row i.
topk_vals, topk_idx = torch.topk(probs, k, dim=-1) # [B, k]
# Apply top-p on the ORIGINAL (non-renormed) probabilities.
# Because the top-k tokens are the k largest in the vocab, their cumulative
# sum in descending order is identical to the full-vocab cumsum for positions
# 0..k-1. This matches the evaluator's valid-mask semantics exactly.
#
# Match evaluator _compute_valid_sampling_mask semantics:
# - Only apply top-p filter when p is strictly in (0, 1).
# - For p <= 0 or p >= 1, keep all top-k tokens (no top-p mask).
# - eps=0.05 tolerance for boundary tokens: cumsum <= p + eps
EPS = 0.05
cumsum = torch.cumsum(topk_vals, dim=-1) # [B, k]
# For out-of-range top_p (e.g. random negatives), use threshold=2.0 so all
# top-k positions pass, matching the evaluator's "skip top-p" behavior.
valid_p = (top_p > 0) & (top_p < 1) # [B]
threshold = torch.where(valid_p, top_p + EPS, torch.full_like(top_p, 2.0))
nucleus = cumsum <= threshold.unsqueeze(-1) # [B, k]
nucleus[:, 0] = True # always keep at least the top token
# Filter and renorm within the top-k candidates
filtered = topk_vals * nucleus.float() # [B, k]
filtered = filtered / filtered.sum(dim=-1, keepdim=True)
# Sample a position within the k-candidate distribution, then map back
# to the original vocabulary index.
pos = flashinfer.sampling.sampling_from_probs(
filtered, deterministic=False
) # [B], values in [0, k)
samples = topk_idx.gather(1, pos.long().unsqueeze(1)).squeeze(1) # [B]
return samples.to(torch.int64)
scrolls · 53 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON