Skip to content

vllm.v1.sample.ops.topk_topp_triton

Combined Top-K and Top-P Triton kernels.

Based on the paper "Qrita: High-performance Top-k and Top-p Algorithm for GPUs using Pivot-based Truncation and Selection" By Park et al. (https://arxiv.org/abs/2602.01518)

Functions:

_apply_topp_split(logits, k, p, mask_value, num_sm)

Split-row top-p pipeline for small batches; masks logits in place.

Source code in vllm/v1/sample/ops/topk_topp_triton.py
def _apply_topp_split(
    logits: torch.Tensor,
    k: torch.Tensor | None,
    p: torch.Tensor,
    mask_value: float,
    num_sm: int,
) -> None:
    """Split-row top-p pipeline for small batches; masks logits in place."""
    ws = _TRITON_SPLIT_CACHE.get(logits.device)
    if ws is None:
        ws = {
            "stats": logits.new_empty((_SPLIT_MAX_BATCH * _SPLIT_MAX_SPLITS, 4)),
            "parts": logits.new_empty(
                (_SPLIT_MAX_BATCH, _SPLIT_ROUNDS, _SPLIT_MAX_SPLITS, _SPLIT_FANOUT, 3)
            ),
        }
        _TRITON_SPLIT_CACHE[logits.device] = ws

    _topp_split_stats(
        logits,
        ws["stats"],
        k,
        p,
        num_sm,
    )
    for round_i in range(_SPLIT_ROUNDS):
        _topp_split_step(
            logits,
            ws["stats"],
            ws["parts"],
            k,
            p,
            round_i,
            num_sm,
        )
    _topp_split_mask(
        logits,
        ws["stats"],
        ws["parts"],
        k,
        p,
        mask_value,
        num_sm,
    )

_topp_sb_combine(PART_BASE, Z, p, lo, hi, best, best_mass, best_minl, best_nmin, S, F, IS_LAST)

Combine one round's partials and advance the search state.

Returns (done, nomask, pivot, dup, numdup, numkeep, lo, hi, best, best_mass, best_minl, best_nmin, sel_fidx, best_sel). best* tracks the tightest fallback: the largest evaluated pivot whose above-mass is >= p, which guarantees the top-p mass invariant if the search fails to converge. Pivots are recomputed from (lo, hi) via the same ladder the sweep used, so partials line up with pivots without storing them. sel_fidx / best_sel report which ladder slot produced dup / updated best this round (-1 if none), so callers can recover the per-slice (minl, nmin) partials behind the boundary value for deterministic duplicate handling.

Source code in vllm/v1/sample/ops/topk_topp_triton.py
@triton.jit
def _topp_sb_combine(
    PART_BASE,
    Z,
    p,
    lo,
    hi,
    best,
    best_mass,
    best_minl,
    best_nmin,
    S: tl.constexpr,
    F: tl.constexpr,
    IS_LAST: tl.constexpr,
):
    """Combine one round's partials and advance the search state.

    Returns (done, nomask, pivot, dup, numdup, numkeep, lo, hi, best,
    best_mass, best_minl, best_nmin, sel_fidx, best_sel). `best*` tracks the
    tightest fallback: the largest evaluated pivot whose above-mass is >= p,
    which guarantees the top-p mass invariant if the search fails to
    converge. Pivots are recomputed from (lo, hi) via the same ladder the
    sweep used, so partials line up with pivots without storing them.
    `sel_fidx` / `best_sel` report which ladder slot produced `dup` /
    updated `best` this round (-1 if none), so callers can recover the
    per-slice (minl, nmin) partials behind the boundary value for
    deterministic duplicate handling.
    """
    fidx = tl.arange(0, F)
    sidx = tl.arange(0, S)
    offs2 = sidx[:, None] * (F * 3) + fidx[None, :] * 3
    mass_s = tl.load(PART_BASE + offs2 + 0)
    minl_s = tl.load(PART_BASE + offs2 + 1)
    nmin_s = tl.load(PART_BASE + offs2 + 2)
    mass = tl.sum(mass_s, axis=0)
    gmin = tl.min(minl_s, axis=0)
    nmin = tl.sum(tl.where(tl.abs(minl_s - gmin[None, :]) < 1e-9, nmin_s, 0.0), axis=0)
    pivs = _topp_sb_ladder(lo, hi, F)

    ok = (mass >= p) & (mass - gmin * nmin < p)
    ok_idx = tl.max(tl.where(ok, fidx, -1).to(tl.float32))
    ge = mass >= p
    ge_idx = tl.max(tl.where(ge, fidx, -1).to(tl.float32))

    best_sel = -1.0
    if ge_idx >= 0:
        bmask = fidx == ge_idx.to(tl.int32)
        bp = tl.sum(tl.where(bmask, pivs, 0.0))
        if bp > best:
            best = bp
            best_mass = tl.sum(tl.where(bmask, mass, 0.0))
            best_minl = tl.sum(tl.where(bmask, gmin, 0.0))
            best_nmin = tl.sum(tl.where(bmask, nmin, 0.0))
            best_sel = ge_idx

    done = 0.0
    nomask = 0.0
    pivot = 0.0
    sel_mass = 0.0
    dup = 1.0
    numdup = 1.0
    numkeep = 1.0
    sel_fidx = -1.0
    if ok_idx >= 0:
        bmask = fidx == ok_idx.to(tl.int32)
        pivot = tl.sum(tl.where(bmask, pivs, 0.0))
        sel_mass = tl.sum(tl.where(bmask, mass, 0.0))
        dup = tl.sum(tl.where(bmask, gmin, 0.0))
        numdup = tl.sum(tl.where(bmask, nmin, 0.0))
        sel_fidx = ok_idx
        done = 1.0
    else:
        new_lo = tl.maximum(lo, tl.max(tl.where(ge, pivs, -float("inf"))))
        new_hi = tl.minimum(hi, tl.min(tl.where(~ge, pivs, float("inf"))))
        if IS_LAST or not (new_hi > new_lo * (1.0 + 1e-7)):
            # Out of rounds or the range collapsed: fall back to the
            # tightest pivot with mass >= p, or keep the whole row if no
            # pivot qualified.
            if best > 0.0:
                pivot = best
                sel_mass = best_mass
                dup = best_minl
                numdup = best_nmin
            else:
                nomask = 1.0
            done = 1.0
        lo = new_lo
        hi = new_hi
    if (done != 0.0) and (nomask == 0.0):
        # Mirror the monolithic guard: a pivot at/above the max logit would
        # mask everything under the strict `>` comparison.
        nomask = tl.where(pivot * Z >= 1.0, 1.0, 0.0)
        # Duplicate (boundary-value) handling, same formula as the
        # monolithic kernel: keep only some of the boundary duplicates.
        numkeep = numdup - tl.cast(tl.maximum(sel_mass - p, 0.0) / dup, tl.uint32).to(
            tl.float32
        )
        numkeep = tl.minimum(tl.maximum(numkeep, 1.0), numdup)
    return (
        done,
        nomask,
        pivot,
        dup,
        numdup,
        numkeep,
        lo,
        hi,
        best,
        best_mass,
        best_minl,
        best_nmin,
        sel_fidx,
        best_sel,
    )

_topp_sb_ladder(lo, hi, F)

Geometric ladder of F pivots inside (lo, hi).

Source code in vllm/v1/sample/ops/topk_topp_triton.py
@triton.jit
def _topp_sb_ladder(lo, hi, F: tl.constexpr):
    """Geometric ladder of F pivots inside (lo, hi)."""
    fidx = tl.arange(0, F).to(tl.float32)
    return lo * tl.exp(tl.log(hi / lo) * (fidx + 1.0) / (F + 1.0))

_topp_sb_row_stats(STATS, row, S)

Combine per-slice partials into row (max, sum_exp, min, n_finite).

Source code in vllm/v1/sample/ops/topk_topp_triton.py
@triton.jit
def _topp_sb_row_stats(STATS, row, S: tl.constexpr):
    """Combine per-slice partials into row (max, sum_exp, min, n_finite)."""
    sidx = tl.arange(0, S)
    sb = STATS + row * S * 4
    m_s = tl.load(sb + sidx * 4 + 0)
    exp_s = tl.load(sb + sidx * 4 + 1)
    mn_s = tl.load(sb + sidx * 4 + 2)
    c_s = tl.load(sb + sidx * 4 + 3)
    M = tl.max(m_s)
    w = tl.where(m_s == -float("inf"), 0.0, tl.exp(m_s - M))
    Z = tl.sum(exp_s * w)
    return M, Z, tl.min(mn_s), tl.sum(c_s)

_update_min_larger_stats(data, above_mask, min_larger, num_min_larger, sentinel)

Update running (min, count) of values above a pivot across tiles.

Tracks the smallest value strictly above a pivot and how many times it occurs. Called once per tile per pivot; the running state is carried across tiles via min_larger / num_min_larger.

Merge rule
  • tile min < running min → replace both
  • tile min == running min → accumulate count
  • tile min > running min → keep running values
Source code in vllm/v1/sample/ops/topk_topp_triton.py
@triton.jit
def _update_min_larger_stats(data, above_mask, min_larger, num_min_larger, sentinel):
    """Update running (min, count) of values above a pivot across tiles.

    Tracks the smallest value strictly above a pivot and how many times
    it occurs.  Called once per tile per pivot; the running state is
    carried across tiles via `min_larger` / `num_min_larger`.

    Merge rule:
      - tile min < running min  → replace both
      - tile min == running min → accumulate count
      - tile min > running min  → keep running values
    """
    tile_min = tl.min(tl.where(above_mask, data, sentinel))
    tile_eq = above_mask & (tl.abs(data - tile_min) < 1e-9)
    tile_cnt = tl.sum(tile_eq)
    is_new = tile_min < min_larger
    is_same = tl.abs(tile_min - min_larger) < 1e-9
    num_min_larger = tl.where(is_new, tile_cnt, num_min_larger + tile_cnt * is_same)
    min_larger = tl.minimum(min_larger, tile_min)
    return min_larger, num_min_larger

apply_top_k_top_p_triton(logits, k, p, mask_value=float('-inf'))

Apply combined top-k and top-p masking using Triton.

Top-k is applied first (by logit value), then top-p is applied to the remaining k values (by probability).

Parameters:

  • logits

    (Tensor) –

    [batch_size, vocab_size] float32 tensor. The returned tensor may alias this input or be a new contiguous tensor for unsupported layouts.

  • k

    (Tensor | None) –

    [batch_size] int32 tensor of top-k values per row, or None to disable top-k

  • p

    (Tensor | None) –

    [batch_size] float32 tensor of top-p values per row (0 to 1), or None to disable top-p

  • mask_value

    (float, default: float('-inf') ) –

    Value for masked positions (default: -inf)

Returns:

  • Tensor –

    The masked logits tensor. It may or may not be modified in-place.

Source code in vllm/v1/sample/ops/topk_topp_triton.py
def apply_top_k_top_p_triton(
    logits: torch.Tensor,
    k: torch.Tensor | None,
    p: torch.Tensor | None,
    mask_value: float = float("-inf"),
) -> torch.Tensor:
    """Apply combined top-k and top-p masking using Triton.

    Top-k is applied first (by logit value), then top-p is applied
    to the remaining k values (by probability).

    Args:
        logits: [batch_size, vocab_size] float32 tensor. The returned tensor
            may alias this input or be a new contiguous tensor for unsupported
            layouts.
        k: [batch_size] int32 tensor of top-k values per row, or None to disable top-k
        p: [batch_size] float32 tensor of top-p values per row (0 to 1),
            or None to disable top-p
        mask_value: Value for masked positions (default: -inf)

    Returns:
        The masked logits tensor. It may or may not be modified in-place.

    """
    assert logits.ndim == 2
    assert logits.dtype == torch.float32
    batch_size, vocab_size = logits.shape
    topk_enabled = k is not None
    topp_enabled = p is not None

    if batch_size == 0 or not (topk_enabled or topp_enabled):
        return logits

    # The Triton kernel supports arbitrary row strides, but it still assumes
    # the vocab dimension is laid out contiguously within each row.
    if logits.stride(1) != 1:
        logits = logits.contiguous()

    if k is not None:
        assert k.ndim == 1 and k.shape[0] == batch_size
        k_ptr = k.to(torch.int32)
    else:
        k_ptr = logits  # Dummy pointer (won't be read)

    if p is not None:
        assert p.ndim == 1 and p.shape[0] == batch_size
        p_ptr = p.to(torch.float32)
    else:
        p_ptr = logits  # Dummy pointer (won't be read)

    num_sm = num_compute_units(logits.device.index)

    # At small batch sizes the monolithic kernel's standalone top-p path is
    # latency-bound (one program per row, serial vocab sweeps), so use the
    # split-row pipeline for rows without an active top-k. Rows with an
    # active top-k still go through the monolithic kernel (its truncation
    # path is already fast). The monolithic's standalone top-p branch then
    # covers only rows whose top-k is a no-op (k < vocab but <= k finite
    # logits, e.g. grammar masks), which the split pipeline skips.
    use_split = (
        topp_enabled and batch_size <= _SPLIT_MAX_BATCH and logits.device.type == "cuda"
    )
    if use_split and not topk_enabled:
        _apply_topp_split(logits, None, p_ptr, mask_value, num_sm)
        return logits

    NUM_PROGRAMS = min(num_sm, batch_size)

    # Cache per-Triton Program buffer on each device.
    buf_key = (logits.device, logits.dtype, vocab_size)
    buffer = _TRITON_BUFFER_CACHE.get(buf_key)
    if buffer is None or buffer.shape[0] < NUM_PROGRAMS:
        size = min(next_power_of_2(NUM_PROGRAMS), num_sm)
        buffer = logits.new_empty((size, vocab_size))
        _TRITON_BUFFER_CACHE[buf_key] = buffer
    if buffer.shape[0] > NUM_PROGRAMS:
        buffer = buffer[:NUM_PROGRAMS]

    # Cache lookup table entries on each device.
    tables = _TRITON_TABLE_CACHE.get(logits.device)
    if tables is None:
        with gpu_sync_allowed():
            normal_cdf_to_sigma_table = logits.new_tensor(_NORMAL_CDF_TO_SIGMA_TABLE)
            percentile_to_std_table = logits.new_tensor(_PERCENTILE_TO_STD_TABLE)
            _TRITON_TABLE_CACHE[logits.device] = (
                normal_cdf_to_sigma_table,
                percentile_to_std_table,
            )
    else:
        normal_cdf_to_sigma_table, percentile_to_std_table = tables

    _topk_topp(
        logits,
        buffer,
        percentile_to_std_table,
        normal_cdf_to_sigma_table,
        k_ptr if topk_enabled else None,
        p_ptr if topp_enabled else None,
        mask_value,
        num_sm,
    )
    if use_split:
        _apply_topp_split(logits, k_ptr, p_ptr, mask_value, num_sm)

    return logits