面向 LLM 采样的免排序 GPU 内核
背景
随着大型语言模型(LLM)词汇量不断增大,分类采样(token 选择)已成为 LLM 推理服务中的一个显著性能瓶颈。FlashInfer 中的采样操作符首次在v0.0.5版本中引入,此后 FlashInfer 团队持续改进其鲁棒性和性能。在这篇博文中,我们将深入探讨 FlashInfer 采样操作符背后的算法和实现细节。
LLM 采样
分类采样是从模型输出概率(覆盖整个词汇表)中选择下一个特定 token 的过程。在实践中,在采样之前会应用过滤以跳过概率可忽略的 token、控制生成行为并强制执行最低概率,例如 Top-P、Top-K 或 Min-P 阈值。
图1:突出显示采样过程的计算时间分解。在 vLLM 1xH100 配置中,我们的内核在所有三个模型上将总体采样时间减少了 50% 以上。
-
Top-K
Top-K 采样在每个生成步骤只保留概率最高的 $K$ 个 token。例如,如果 $K=50$,模型将忽略前 50 个最有可能候选项之外的所有 token。
-
Top-P 则保留最小的一组 token,其累积概率刚好超过阈值 $P$。例如,如果 $P=0.9$,您按降序累积 token 概率,直到它们的总和至少为 0.9。
-
Min-P 过滤掉所有低于最小阈值 $p_\text{base} \times p_\text{max}$ 的 token,其中 $p_\text{base}$ 是参数,$p_\text{max}$ 是输入中最大的概率。这有助于消除极不可能的 token,同时仍尊重顶级候选项之间的相对差异。
在实践中,Top-K 和 Top-P 过滤的组合很受欢迎,并被用作 LLM 采样的标准设置。这允许对生成过程进行更精细的控制。例如,如果我们首先使用 Top-K 过滤,我们首先将 token 集限制为 Top-K 最高概率,然后应用 Top-P 截止来过滤这些 $K$ 个 token 中的尾部部分。1
这些采样器的一个 PyTorch 实现可能如下所示
# vllm/vllm/model_executor/layers/sampler.py
def _apply_top_k_top_p(
logits: torch.Tensor,
p: torch.Tensor,
k: torch.Tensor,
) -> torch.Tensor:
logits_sort, logits_idx = logits.sort(dim=-1, descending=False)
# Apply top-k.
top_k_mask = logits_sort.size(1) - k.to(torch.long)
# Get all the top_k values.
top_k_mask = logits_sort.gather(1, top_k_mask.unsqueeze(dim=1))
top_k_mask = logits_sort < top_k_mask
logits_sort.masked_fill_(top_k_mask, -float("inf"))
# Apply top-p.
probs_sort = logits_sort.softmax(dim=-1)
probs_sum = probs_sort.cumsum(dim=-1)
top_p_mask = probs_sum <= 1 - p.unsqueeze(dim=1)
# at least one
top_p_mask[:, -1] = False
logits_sort.masked_fill_(top_p_mask, -float("inf"))
# Re-sort the probabilities.
src = torch.arange(logits_idx.shape[-1],
device=logits_idx.device).expand_as(logits_idx)
logits_idx_inv = torch.empty_like(logits_idx).scatter_(dim=-1,
index=logits_idx,
src=src)
logits = torch.gather(logits_sort, dim=-1, index=logits_idx_inv)
return logits
这段代码使用了排序、累积求和和掩码的组合。虽然很容易理解,但它会导致性能瓶颈,尤其是对于大词汇量,因为排序的开销巨大。
在 FlashInfer 中,我们展示了带过滤的采样可以以免排序的方式完成,我们引入了双枢轴拒绝采样算法并设计了融合采样内核模板,以充分利用 GPU 的并行计算能力,最终实现对数时间复杂度(在最坏情况下)。在这篇博文中,我们将向您介绍我们如何开发这种算法,整合了逆向采样、拒绝采样的思想以及具有收敛理论保证的最终算法版本。
算法
逆变换采样
图2:逆变换采样。此动画说明了每块(per-block)的过程,在实践中,工作负载按块执行。
我们首先实现一个基本的采样内核,它纯粹根据 token 的概率来选择 token,特别是在 GPU 并行计算环境中。
该方法是逆变换采样,它根据给定概率分布的累积分布函数(CDF)从中抽取样本。对于 token 采样过程,CDF 将是 token 概率的前缀和。算法如下所示:
- 从 $U\sim \text{Unif}(0,1)$ 中抽取一个随机数 $u$。
- 计算每个采样 token $j$ 的前缀和(CDF),概率为 $p_j$:$F_j=\sum^{j}_{i=1}p_i$。
- 定位 token $k$,使得 $F_{k-1} \leq u < F_k$ 作为结果。
NVIDIA 的 CUB 库(现为 CCCL 的一部分)提供了用于并行计算的高效原语,我们利用 reduce 和 scan 原语来计算前缀和。我们为每个概率分布使用一个 threadblock,对于批量采样,我们并行启动多个 threadblock。块级 reduce/scan 原语可以应用于一个元素块(BLOCK_SIZE = NUM_THREADS * NUM_ELEMENTS_PER_THREADS,例如,对于浮点输入,1024 * 4),对于大于 BLOCK_SIZE 的词汇量,我们将词汇表分成多个块并依次对每个块应用相同的过程。
- 初始化一个运行总计 $\texttt{a}=0$。计算每个块的概率总和 $\texttt{a_local}$。如果 $\texttt{a} + \texttt{a_local}> u$,则采样的 token 位于此块中。
- 如果不是,我们将 $\texttt{a_local}$ 添加到 $\texttt{a}$ 并转到下一个块。
- 一旦知道正确的块,我们就会对其中的 token 执行前缀和以精确定位 token 索引。
我们使用 BlockReduce 和 BlockScan 进行每块部分求和和前缀和,以及 AdjacentDifference 来定位 token 索引。在实践中,我们使用提前停止来终止逆变换采样过程,当累积概率超过随机数 $u$ 时,这样我们就不需要为每一轮遍历整个词汇表。
拒绝采样
图3:Top-P 拒绝采样。此动画说明了每块(per-block)的过程,在实践中,工作负载按块执行。
对于更高级的策略,例如 Top-P 采样,我们使用拒绝采样来限制可以选择的 token。拒绝采样通过将随机样本与阈值进行比较并丢弃不符合要求的样本来从目标分布中抽取样本。
以 Top-P 过滤下的采样内核为例,以下是简化后的发生情况:
- 初始化枢轴为 $0$,因此最初所有 token 都被考虑在内。
- 执行一次逆变换采样过程,但忽略概率低于当前枢轴的 token。采样一个 token 后,将枢轴更新为该 token 的概率。
- 计算仍在该枢轴之上的 token 之间的剩余概率 $\texttt{q}$
- 如果 $\texttt{q}$ 仍然大于或等于 $\texttt{top_p}$,则需要另一轮来进一步提高枢轴并拒绝更多 token。
- 否则,如果它低于 $\texttt{top_p}$,我们确定采样的 token 并标记成功。
- 重复直到成功。
整个算法可以在单个融合内核中实现,它对 Top-K 和其他过滤策略的工作方式类似,只是我们将检查超过枢轴的 token 数量是否符合 $\texttt{top_k}$ 或 $\texttt{min_p}$。
在实践中,我们发现返回采样 token 所需的轮数通常很少。与朴素的 PyTorch 实现相比,它提供了显着的加速,因为我们避免了排序和多次遍历词汇表,以及多次内核启动开销。
双枢轴拒绝采样
虽然这种拒绝采样方法在大多数情况下简单高效,但它有一些局限性。对于获取采样 token 所需的轮数没有理论保证。这可能导致不同概率分布的采样时间不同,进而导致 LLM 推理服务期间的 inter-token 延迟不一致。这种可变性可能会影响服务系统的可预测性和可靠性。
为了解决这个问题,在 FlashInfer v0.2.3 中,我们引入了一种名为双枢轴拒绝采样的新算法,它使用两个枢轴在拒绝采样中实现更快收敛。算法如下:
- 设 $f$ 是一个检查概率值是否有效的函数:如果有效,则 $f(x)=1$,否则为 $0$。
- 初始化 $\textrm{low} \leftarrow 0$ 和 $\textrm{high} \leftarrow \max_i(p_i)$ 作为初始范围,保证 $f(\textrm{low})=0$ 和 $f(\textrm{high})=1$。
- 使用逆变换采样对范围 $(\textrm{low}, \infty)$ 中的概率值进行采样。
- 假设 $j$ 是采样的 token,设 $\textrm{pivot}_1\leftarrow p_j$,$\textrm{pivot}_2\leftarrow \frac{\textrm{pivot}_1+\textrm{high}}{2}$。
- 如果 $f(\textrm{pivot}_1)=1$,我们接受采样的 token 并返回 $j$。
- 如果 $f(\textrm{pivot}_1)=0$,$f(\textrm{pivot}_2)=1$,我们将 $\textrm{pivot}_1$ 设置为新的 $\textrm{low}$,并将 $\textrm{pivot}_2$ 设置为新的 $\textrm{high}$。
- 如果 $f(\textrm{pivot}_1)=0$,$f(\textrm{pivot}_2)=0$,我们将 $\textrm{pivot}_2$ 设置为新的 $\textrm{low}$。
- 重复步骤 3 和 4 直到成功。
图4:双枢轴拒绝采样中从第 i 轮到第 i+1 轮的过渡,我们要么接受采样的 token(情况1),要么将范围缩小至少一半(情况2和3)。
图4显示了双枢轴拒绝采样中从第 i 轮到第 i+1 轮的过渡,在每一轮中,如果采样的 token 被接受,我们返回 token,否则,新范围的跨度是 $\frac{\text{high}-\text{pivot}_1}{2} < \frac{\text{high}-\text{low}}{2}$,这至少是前一个范围的一半。因此,保证轮数是 $O(\log(1/\epsilon))$,其中 $\epsilon$ 是浮点表示中可能的最小正值。
拒绝采样器正确性的理论证明
在本节中,我们提供拒绝采样器正确性的理论证明,我们以 top-k 采样为例,其他采样器可以用类似的方式证明。
命名法
| 符号 | 含义 |
|---|---|
| $p_i > 0$ | 项目 $i$ 的未归一化分数(未归一化概率质量) |
| $T = \operatorname{Top}k = {i_1,\dots,i_k}$ | k 个最大分数的索引 |
| $Z = \sum_{j \in T} p_j$ | 前 k 个项目的总质量 |
| $\tau$ | 当前的枢轴(阈值)值 |
定理
算法输出每个 top-k 索引 $j \in T$ 的概率为
\[\Pr[\text{output}=j] \;=\; \frac{p_j}{Z},\]即恰好是首先丢弃所有非 top-k 项目然后在 top-k 集合中进行分类采样所获得的分布。
证明
固定任意枢轴 $\tau < \min_{j \in T} p_j$(在每一步都成立,因为 $\tau$ 总是取自被拒绝的非 top-k 项目)。定义
\[Q_j(\tau) \;=\; \Pr[\text{algorithm eventually returns } j \mid \text{current pivot } \tau], \quad j \in T .\]其中
\[S(\tau) \;=\; \sum_{m : p_m > \tau} p_m \;=\; Z \;+\; W(\tau),\qquad W(\tau) \;=\!\!\!\! \sum_{r \notin T,\, p_r > \tau}\!\!\!\! p_r ,\]其中 $S(\tau)$ 是所有大于 $\tau$ 的 token 概率的总和,$W(\tau)$ 是仍高于阈值的“坏”项目(非 top-k 项目)的剩余质量。
下一次抽取遵循
\[\Pr[i \mid \tau] \;=\; \frac{p_i}{S(\tau)}.\]因此
\[Q_j(\tau) \;=\; \underbrace{\frac{p_j}{S(\tau)}}_{\text{立即接受}} \;+\; \sum_{\substack{r \notin T \\ p_r > \tau}} \underbrace{\frac{p_r}{S(\tau)}}_{\text{抽取 } r}\; \underbrace{Q_j\!\bigl(p_r\bigr)}_{\text{枢轴变为 } p_r} \tag{★}\]我们证明以下公式是 (★) 的有效解
\[\boxed{\,Q_j(\tau) \;=\; \dfrac{p_j}{Z}\,} \qquad\text{对于每个 }\tau < \min_{j \in T} p_j .\]我们可以通过将其代入 (★) 来验证解
\[\begin{aligned} \text{RHS} &= \frac{p_j}{S(\tau)} \;+\; \frac{p_j}{Z} \frac{W(\tau)}{S(\tau)} \\ &= \frac{p_j}{S(\tau)}\!\Bigl(1+\frac{W(\tau)}{Z}\Bigr) \\ &= \frac{p_j}{Z} \frac{Z+W(\tau)}{S(\tau)} \\ &= \frac{p_j}{Z}, \end{aligned}\]因为 $S(\tau) = Z + W(\tau)$。因此,所声称的形式满足递推关系,所以 $Q_j(\tau) \equiv p_j/Z$。
现在让我们证明解是唯一的。假设存在另一个解 $Q_j’(\tau)$ 满足 (★),我们定义 $\Delta_j(\tau) = Q_j(\tau) - Q_j’(\tau)$,我们有
\[\Delta_j(\tau) = \sum_{\substack{r \notin T \\ p_r > \tau}} \frac{p_r}{S(\tau)} \Delta_j(p_r)\]系数的总和 $\sum_{\substack{r \notin T \ p_r > \tau}} \frac{p_r}{S(\tau)} = \frac{W(\tau)}{S(\tau)}$,它满足
\[0 \leq \frac{W(\tau)}{S(\tau)} < 1\]假设 $\tau^*$ 是 $| \Delta_j(\tau)|$ 达到最大值时的枢轴,如果它是正的,我们有
\[|\Delta_j(\tau^*)| \leq \sum_{\substack{r \notin T \\ p_r > \tau^*}} \frac{p_r}{S(\tau)} |\Delta_j(p_r)| \leq \sum_{\substack{r \notin T \\ p_r > \tau^*}} \frac{p_r}{S(\tau)} |\Delta_j(\tau^*)| = \frac{W(\tau^*)}{S(\tau^*)} |\Delta_j(\tau^*)| < |\Delta_j(\tau^*)|\]这导致矛盾,这意味着 $\Delta_j(\tau^*) = 0$,我们的解是唯一的。
算法从 $\tau = 0$ 开始;因此
\[\Pr[\text{output}=j] = Q_j(0) = \frac{p_j}{Z},\]恰好是所需的 top-k 分类分布。
评估
我们的评估表明,与传统的基于排序的实现相比,FlashInfer 的采样内核在内核级延迟和端到端吞吐量方面都提供了显着的改进。
图5:不同引擎内核的吞吐量比较。
图6:采样延迟随批量大小的增长。
社区采用和其他应用
FlashInfer 采样内核已在主要 LLM 框架中得到广泛采用,包括 MLC-LLM、sglang 和 vLLM。社区通过反馈和错误报告的积极参与对完善和改进我们的实现起到了重要作用。
除了 token 采样之外,拒绝采样算法在 LLM 推理优化的其他领域也证明了其价值。类似的算法也可以应用于推测解码验证,例如 链式推测采样和 树形推测采样。最近的创新,如 Twilight,通过将 top-p 采样与稀疏注意力以统一的方法结合起来,进一步推进了这一领域。
实现细节
虽然算法在理论上很优雅,但在 GPU 内核中高效实现它需要仔细注意细节,尤其是在逆变换采样中的 token 选择逻辑。一个关键挑战在于用于定位采样 token 的并行前缀和操作。由于浮点运算的非结合性和非交换性,即使输入是非负数,并行前缀和也不能保证单调输出。如果处理不当,这可能导致无效的 token 生成。必须特别注意确保采样实现中的数值稳定性和正确性(我们在得到正确结果之前犯了很多错误)。
要详细了解我们的实现以及我们如何应对这些挑战,您可以浏览我们的源代码。此外,FlashInfer 提供了一套用于概率截止和重新归一化的综合 API,例如 top_p_renorm_probs 和 top_k_renorm_probs,可以灵活组合多个采样过滤器。这些工具使开发人员能够构建针对其特定需求量身定制的复杂采样策略。
致谢
这篇博文由 Shanli Xing撰写,我们感谢 flashinfer 团队对 flashinfer.sampling 模块的贡献
- Zihao Ye:CUDA 中采样内核的设计和实现。
- Bohan Hou:TVM 中采样内核的设计和实现。
- Shanli Xing:CUDA 中 min-p 采样内核的设计和实现。
- Tianqi Chen:提出了用于 top-p 的拒绝采样的想法。

评论