跳转至

解码与采样

本页讨论采样:运行时如何把 logits 变成输出 token。Speculative decoding 见 Speculative Decoding。

采样

采样之所以重要,是因为仅有 next-token 概率还不足以决定交互式模型的行为。系统仍然必须决定输出是更保守还是更多样,是更稳定还是更探索,不同采样规则正是在这些目标之间暴露不同权衡。

Temperature

温度 \(T\) 控制概率分布的尖锐度。给定模型 logits \(z_i\),温度采样会在 softmax 前对其重缩放:

\[ p_i=\frac{e^{z_i / T}}{\sum_j e^{z_j / T}} \]
  • \(T < 1\) 时分布更尖锐,随机性更低;\(T > 1\) 时分布更平坦,探索性更高。
  • 温度不会改变 logits 的排序,因为除以正数不会改变相对顺序。如果 top-k = 1,也就是 greedy 采样,温度基本不起作用。
  • 当 \(T \rightarrow 0\) 时,softmax 分布会收敛到 greedy。
  • 当 \(T \rightarrow \infty\) 时,分布趋于均匀。

Top-k Sampling

在计算概率之后,按概率排序,仅保留概率最高的 \(k\) 个 token,然后对 \(p_i^{\prime}\) 重新归一化:

\[ \begin{aligned} S_k & = \text{top-k tokens by } p_i \\ p_i^{\prime} & = \begin{cases} \frac{p_i}{\sum_{j \in S_k} p_j} & i \in S_k \\ 0 & \text{otherwise} \end{cases} \end{aligned} \]

然后从该截断分布中采样。

  • 较小的 \(k\) 会让输出更确定、更保守。
  • 较大的 \(k\) 会提升多样性,但也会增加偏题 token 的概率。

Top-p(Nucleus)Sampling

与固定 \(k\) 不同,top-p 选择累计概率质量达到阈值 \(p\) 的最小 token 集合:

\[ \begin{aligned} S_p & = \left\{i : \sum_{j \in S_p} p_j \geq p \right\} \\ p_i^{\prime} & = \begin{cases} \frac{p_i}{\sum_{j \in S_p} p_j} & i \in S_p \\ 0 & \text{otherwise} \end{cases} \end{aligned} \]

它与 top-k 的关键区别在于会根据模型不确定性自适应调整候选集大小。

  • 如果模型很确定,保留的 token 很少。
  • 如果模型不确定,候选集会自动扩大。

Min-p

Min-p 会过滤掉相对 top-1 来说过于不可能的 token。它不固定 \(k\) 或 \(p\),而是保留那些概率至少达到 top-1 token 一定比例的 token:

\[ S_{\min }=\left\{i: p_i \geq \min -\mathrm{p} \times p_{\max }\right\} \]

如果 min-p = 0.1,则任何概率至少达到最可能 token 的 10% 的 token 都会被保留。随后归一化:

\[ p_i^{\prime} = \frac{p_i}{\sum_{j \in S_{\min}} p_j} \]
  • 和 top-p 一样,min-p 会随上下文自适应。
  • 在分布较平坦时,它更容易保留一些虽然排名较低但仍合理的 token。
  • 对长尾词表通常也更数值稳定。

组合方式

常见的组合顺序是:

Temperature -> Softmax -> (top-k / top-p / min-p) -> Renormalize -> Sample

当 top-k 和 top-p 同时启用时,最终保留的是二者交集。

示例代码

代码片段来自 Omniinfer Sampler。

Top-k 和 Top-p 实现
def apply_top_k_top_p(
    logits_or_prob: torch.Tensor,
    k: Optional[torch.Tensor],
    p: Optional[torch.Tensor],
    is_logits: bool,
) -> torch.Tensor:
    if p is None:
        if k is not None:
            logits_or_prob = apply_top_k_only(logits_or_prob, k, is_logits)
        if is_logits:
            probs = logits_or_prob.softmax(dim=-1, dtype=torch.float32)
        else:
            probs = logits_or_prob / logits_or_prob.sum(dim=-1, keepdim=True)
        return probs, None

    logits_or_prob_sort, logits_or_prob_idx = logits_or_prob.sort(dim=-1, descending=False)

    if k is not None:
        # Apply top-k.
        top_k_mask = logits_or_prob_sort.size(1) - k.to(torch.long)  # shape: B
        # Get all the top_k values.
        top_k_mask = logits_or_prob_sort.gather(1, top_k_mask.unsqueeze(dim=1))
        top_k_mask = logits_or_prob_sort < top_k_mask
        logits_or_prob_sort.masked_fill_(top_k_mask, -float("inf") if is_logits else 0)

    # Apply top-p.
    if is_logits:
        probs_sort = logits_or_prob_sort.softmax(dim=-1)
    else:
        probs_sort = logits_or_prob_sort / logits_or_prob_sort.sum(dim=-1, keepdim=True)
    probs_sum = torch.cumsum(probs_sort, dim=-1, out=probs_sort)
    top_p_mask = probs_sum <= 1 - p.unsqueeze(dim=1)
    # at least one
    top_p_mask[:, -1] = False
    probs_sort.masked_fill_(top_p_mask, 0)
    probs = probs_sort / probs_sort.sum(dim=-1, keepdim=True)
    return probs, logits_or_prob_idx
def apply_top_k_only(
    logits_or_prob: torch.Tensor,
    k: torch.Tensor,
    is_logits: bool,
) -> torch.Tensor:
    """
    Apply top-k mask to the logits.

    This implementation doesn't involve sorting the entire vocab.

    The logits tensor may be updated in-place.
    """
    no_top_k_mask = k == logits_or_prob.shape[1]
    # Set non-top-k rows to 1 so that we can gather.
    k = k.masked_fill(no_top_k_mask, 1)
    max_top_k = k.max()
    # topk.values tensor has shape [batch_size, max_top_k].
    # Convert top k to 0-based index in range [0, max_top_k).
    k_index = k.sub_(1).unsqueeze(1)
    top_k_mask = logits_or_prob.topk(max_top_k, dim=1).values.gather(1, k_index.long())
    # Handle non-topk rows.
    top_k_mask.masked_fill_(no_top_k_mask.unsqueeze(1), -float("inf"))
    logits_or_prob.masked_fill_(
        logits_or_prob < top_k_mask,
        -float("inf") if is_logits else float(0),
    )
    return logits_or_prob