Skip to content
KernelIndex
Search⌘K

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