拒绝采样
1 投机解码简介
为了提升大语言模型在推理时的解码效率,研究者们想出一种投机解码(Speculative Decoding)方式:
- 令我们想要采样的大模型的采样分布为 $p(x)$,当给定 token $x_ {<t}$ 后,采样 $x_ t$ 的条件分布记为 $p(\cdot \vert x_ {<t})$。
- 随着语言模型越来越大,推理时采样 token $x_ t$ 越来越耗时。因此,研究者希望找到另一个采样比较快的模型 $q(x)$,使用 $q$ 来对 $x_ t$ 进行采样:$x_ t \sim q(\cdot \vert x_ {<t})$。
在给定 $x_ {<t}$ 之后,可以使用 $q$ 采样 $x_ t \sim q(\cdot\vert x_ {<t})$, $x_ {t+1} \sim q(\cdot \vert x_ {<t+1})$, …, $x_ {t+n} \sim q(\cdot \vert x_ {<t+n})$。在拿到使用 $q$ 采样出来的 $n+1$ 个 token $x_ t, x_ {t+1}, \dots, x_ {t+n}$ 之后,再使用目标模型 $p$ 来对这些 token 进行检查,接受符合条件的 token,拒绝不符合条件的 token 并重新采样。
假设最终接受了 $x_ t, \dots, x_ {t+k}$ 一般来说,使用 $q$ 采样 $x_ t, x_ {t+1}, \dots, x_ {t+n}$ 的时间 + 使用 $p$ 检查这些 tokens 的时间不能比使用 $p$ 采样 $x_ t, \dots, x_ {t+k}$ 的时间长,否则这个过程就无法起到加速采样的效果了。
当模型 $p$ 接受 $x_ {t+k}$ 并拒绝 $x_ {t+k+1}$ 之后,一个很自然的想法是基于 $x_ {< t+k+1}$ 使用目标分布 $p$ 去采样 $x_ {t+k+1} \sim p(\cdot \vert x_ {< t+k+1})$:
- 以 multi-token prediction(MTP)为例子,当使用更快的 MTP 头采样出 $x_ t, x_ {t+1}, \dots, x_ {t+n}$ 后,为了验证这些新 token,会使用整个大语言模型计算 $p(x_ t \vert x_ {<t}), \dots, p(x_ {t+n} \vert x_ {< t+n})$。当然,这里计算 $p(x_ t \vert x_ {<t}), \dots, p(x_ {t+n} \vert x_ {< t+n})$ 只需要将 $x_ {<t+n}$ 作为 prefill 并行计算一次,所以所需时间不会比往前推理一个 token 多太多。而且使用 MTP 投采样的效率比使用 $p$ 快很多,因此只要最终接受 token 数 $>2$,那么这个投机采样策略就不亏。假设验证后决定接受 $x_ t, \dots, x_ {t+k}$ 并拒绝 $x_ {t+k+1}$,因为验证的时候可以计算出 $p(\cdot \vert x_ {<t+k+1})$,因此为了利用好已有的分布,一个很自然的想法就是直接使用 $p(\cdot \vert x_ {<t+k+1})$ 采样 $x_ {t+k+1}$。
然而,在拒绝后直接采样 $x_ {t+k+1} \sim p(\cdot \vert x_ {< t+k+1})$ 会使分布发生偏移,即我们采样的 $(x_ t, \dots, x_ {t+k+1}) \nsim p((x_ t, \dots, x_ {t+k+1}) \vert x_ {<t})$。这一点的论证可以看后面的分析。因此,我们使用拒绝采样的方式来进行矫正,使得我们采样到的 $(x_ t, \dots, x_ {t+k+1})$ 服从分布 $p((x_ t, \dots, x_ {t+k+1}) \vert x_ {<t})$。
2 拒绝采样做法
拒绝采样的目的是在加速采样的同时,保证最终的输出严格服从目标模型的概率分布。为方便起见,我们考虑一步采样过程,用 $p(x)$ 表示我们希望采样的目标分布,用 $q(x)$ 表示我们用来加速采样的代理分布。假设代理分布采样出来的 token 是 $x$,那么拒绝采样的做法是:以 $P_ \mathrm{acc}(x)$ 的概率接受这个 token。即:
- 从 $[0,1]$ 上的均匀分布中采样 $u \sim \mathrm{Uniform}(0,1)$。
- 如果 $u \le P_ \mathrm{acc}(x)$ 则接受该 token,然后继续检查下一个 token;如果 $u > P_ \mathrm{acc}(x)$ 则拒绝该 token。
这里选择的 $P_ \mathrm{acc}$ 为:
\[P_ \mathrm{acc}(x) = \min \left( 1, \frac{p(x)}{q(x)} \right).\]直观的理解是:
- 如果 $P_ \mathrm{acc}(x) = 1$,则 $p(x) \ge q(x)$。此时虽然 $x$ 是从 $q$ 中采样出来的,但是目标模型 $p$ 比代理模型 $q$ 更加认可这个 token,因此选择接受该 token、。
- 如果 $P_ \mathrm{acc}(x) < 1$,则 $p(x) < q(x)$。此时代理模型生成的这个 token 的概率相对于目标模型 $q$ 来说过高,属于过度采样。因此,为了削减代理模型对该 token 的过度采样,我们希望以 $\frac{p(x)}{q(x)}$ 的概率来接受 $x$。直观上看就是强行将采样到 $x$ 的概率设置为 $q(x) \cdot \frac{p(x)}{q(x)} = p(x)$,与目标模型采样到 $x$ 的概率对齐。
如果模型拒绝了 token $x$ 之后要怎么处理?前面第 1 节提到的直观做法是直接从目标分布 $p$ 中采样一个 token,但这样会导致采样分布发生偏差。拒绝采样的做法是:从修正后的残差分布中采样新的 token,即从分布
\[p_\mathrm{res}(y) = \frac{\operatorname{ReLU}(p(y) - q(y))}{\sum_ {z \in V} \operatorname{ReLU}(p(y)-q(y))},\]其中 $V$ 为 vocabulary set,$\operatorname{ReLU}(x) = \max (0, x)$ 为 relu 函数。上述残差分布补偿了目标模型相比于代理模型少采样到的那些 token。直观地说就是:
- 当 $q(y) > p(y)$,则代理模型采样的时候已经过度生成(虽然不一定真的采样到了,但是概率更大)了 $y$,于是不再补偿,$p_\mathrm{res}(y) = 0$。
- 当 $q(y) < p(y)$,则代理模型采样的时候对 token $y$ 生成得不够,需要增加 token $y$ 的采样概率。并且使用代理模型采样的时候概率与目标模型差距越大的 token,在拒绝后的重新采样中为它补偿的概率就越多。
3 拒绝采样矫正分布验证
这一节我们来验证,使用拒绝采样确实可以保证我们整个采样过程的分布与目标模型的分布 $p$ 一致。
固定当前上下文 $h$,我们为简单,记 $p(y) = P_ \mathrm{target}(y\vert h)$ 为目标模型的采样分布,记 $q(y) = P_ \mathrm{proxy}(y \vert h)$ 为代理模型的采样分布,记 vocabulary set 为 $V$。
对于拒绝采样,我们来计算 “代理 token 被接受且输出 $y$ 的概率”。要通过 “接受路径” 输出 $y$,那么必须有:(1)代理模型采样得到 $X=y$;(2)验证时接受 $y$。于是有:
\[P(Y=y,\mathrm{accept}) = q(y) \cdot \min \left( 1, \frac{p(y)}{q(y)} \right) = \min \left( q(y), p(y) \right).\]于是总的接受概率为:
\[P(\mathrm{accept}) = \sum_ {z \in V} P(Y=y,\mathrm{accept}) = \sum_ {z \in V} \min \left( q(y), p(y) \right).\]总的拒绝概率为:
\[\begin{align*} P(\mathrm{reject}) &= 1 - P(\mathrm{accept}) = 1 - \sum_ {z \in V} \min \left( q(y), p(y) \right) \\ &= \sum_ {z \in V} p(z) - \sum_ {z \in V} \min \left( q(y), p(y) \right) \\ &= \sum_ {z \in V} \left[ p(z) - \min \left( q(y), p(y) \right) \right] \\ &= \sum_ {z \in V} \operatorname{ReLU}(p(z)- q(z)), \end{align*}\]其中最后一个等式是因为 $a - \min (a,b) = \operatorname{ReLU}(a-b)$ 对任意实数 $a,b$ 均成立。
于是残差分布可以写为:
\[p_ \mathrm{res}(y) = \frac{\operatorname{ReLU}(p(y)-q(y))}{P(\mathrm{reject})}.\]从而有:
\[\begin{align*} P(Y=y) &= P(Y=y,\mathrm{accept}) + P(Y=y,\mathrm{reject}) \\ &= P(Y=y,\mathrm{accept}) + P(Y=y \vert \mathrm{reject}) \cdot P(\mathrm{reject}) \\ &= \min (p(y), q(y)) + p_ \mathrm{res}(y) \cdot P(\mathrm{reject}) \\ &= \min (p(y), q(y)) + \operatorname{ReLU}(p(y)-q(y)). \end{align*}\]分情况讨论:
- 当 $p(y) \ge q(y)$,此时 $\min (p(y), q(y)) = q(y)$,$\operatorname{ReLU}(p(y)-q(y)) = p(y)-q(y)$,因此 $P(Y=y) = q(y) + p(y) - q(y) = p(y)$。
- 当 $p(y) < q(y)$,此时 $\min (p(y), q(y)) = p(y)$,$\operatorname{ReLU}(p(y)-q(y)) = 0$,因此 $P(Y=y) = p(y)$。
综上,$P(Y=y) = p(y)$,故拒绝采样得到的当前 token $y$ 的分布与目标模型的分布一致。
前面考虑的是单 token 投机,下面我们考虑多 token 投机。假设目标模型的联合分布为:
\[p(x_ {1:T}) = \prod_ {t=1}^T p(x_ t \vert x_ {<t}).\]上述假设是与当前大语言模型 next-token perdition 的模式一致的。在投机解码的时候,对于每个位置 $t$,都是基于当前已经确认的前缀 token $x_ {<t}$ 进行解码的,因此,假设从 $t$ 处开始投机解码,代理模型满足:
\[q(x_ {t:T}) = \prod_ {j = t}^T q(x_ j \vert x_ {<j}).\]因此我们有:
\[\begin{align*} p(x_ {1:T}) &= \prod_ {j=1}^T p(x_ j \vert x_ {<j}) \\ &= \prod_ {j=1}^{t-1} p(x_ j \vert x_ {<j}) \cdot \prod_ {j=t}^T p(x_ j \vert x_ {<j}). \end{align*}\]由于前面单 token 投机中证明了对 token $t$ 使用投机时,有 $p(x_ t \vert x_ {<t}) = q(x_ t \vert x_ {<t})$,于是
\[\prod_ {j=t}^T p(x_ j \vert x_ {<j}) = \prod_ {j=t}^T q(x_ j \vert x_ {<j}),\]故
\[p(x_ {1:T}) = \prod_ {j=1}^{t-1} p(x_ j \vert x_ {<j}) \cdot \prod_ {j=t}^T q(x_ j \vert x_ {<j}) .\]上式中等号右边正好是我们在 $t$ 时刻开始多步投机所对应的采样分布,于是证明完成。
4 拒绝采样分布求解
前一节验证了,拒绝采样可以保证整个采样过程的分布与目标模型分布一致。但那是马后炮的验证,假设现在没有拒绝采样的分布矫正机制,别人还是在拒绝后使用 $p(x)$ 进行采样。那么当我们意识到这样做会导致分布偏移之后,我们要怎么样去求解 “在拒绝该token后,到底应该在什么分布上重新采样才能保证分布不发生偏移?”
第 3 节末尾关于多步投机的分析已经告诉我们,只要我们给出单步时的结果,就能自然地使多步投机也保持分布。因此我们考虑单步投机时拒绝后的采样分布求解。一样地,假设目标分布是 $p(y)$,代理模型分布是 $q(y)$,整个拒绝采样过程为:
- 先采样 token $y \sim q(\cdot)$
- 然后根据某个未知的规则 $r(\cdot)$ 判断是否拒绝该 token:采样 $u \sim \operatorname{Uniform}(0,1)$,若 $u < r(y)$ 则接受 token $y$,否则拒绝 token $y$
- 拒绝 token $y$ 之后,根据某一分布 $a(\cdot)$ 采样一个新的 token。
我们的目标是找到规则 $r$ 以及分布 $a$ 使得拒绝采样对应的分布等于目标分布 $p$。与第 3 节的思路类似,我们有:
\[P(X=y, \mathrm{accept}) = q(y) r(y),\]于是接受的总概率为:
\[P(\mathrm{accept}) = \sum_ {z \in V} q(z) r(z),\]故拒绝的总概率为:
\[\begin{align*} P(\mathrm{reject}) &= 1 - \sum_ {z \in V} q(z) r(z) \end{align*}\]从而拒绝采样采样到 token $y$ 的概率为:
\[\begin{align*} P(X=y) &= P(X=y, \mathrm{accept}) + P(X=y,\mathrm{reject}) \\ &= P(X=y, \mathrm{accept}) + P(X=y \vert \mathrm{reject}) \cdot P(\mathrm{reject}) \\ &= q(y)r(y) + a(y)P(\mathrm{reject}). \end{align*}\]我们的目标是 $P(X=y) = p(y)$,故我们需要求解 $q(y)r(y) + a(y)P(\mathrm{reject}) = p(y)$,移项整理得:
\[a(y) = \frac{p(y) - q(y)r(y)}{P(\mathrm{reject})} = \frac{p(y) - q(y)r(y)}{1 - \sum_ {z \in V} q(z) r(z)}.\]如果 $a(\cdot)$ 是一个分数,它需要满足:
- $a(y) \ge 0$ for all $y \in V$.
- $\sum_ {y \in V} a(y) = 1$.
求和得:
\[\sum_ {y \in V} a(y) = \frac{\sum_ {y \in V} p(y) - q(y)r(y)}{1 - \sum_ {z \in V} q(z) r(z)} = \frac{1 -\sum_ { y \in V} q(y)r(y)}{1 - \sum_ {z \in V} q(z) r(z)} = 1,\]因此我们只需要满足 $a(y) \ge 0$ 的条件即可,即 $p(y) - q(y) r(y) \ge 0$,于是有:
\[r(y) \le \frac{p(y)}{q(y)}.\]到此处我们已经解出我们想要的 $r(\cdot), a(\cdot)$ 了,有无数组解。我们只需要选:
- 任意满足 $r(y) \le \frac{p(y)}{q(y)}$ 的规则 $r(\cdot)$
- 给定 $r$ 之后,选拒绝后的采样分布为:$a(y) = \frac{p(y) - q(y)r(y)}{1 - \sum_ {z \in V} q(z) r(z)}$.
但是,我们做投机采样本身就是希望接受率足够高,因此我们希望 $r(y)$ 越大越好,而同时接受概率 $r(y)$ 要满足 $r(y) \le 1$,因此我们取 $r(y) = \min \left( 1, \frac{p(y)}{q(y)}\right)$。这正好就是拒绝采样的策略。将 $r$ 带入易得 $a(y) = \frac{\operatorname{ReLU}(p(y) - q(y))}{\sum_ {z \in V} \operatorname{ReLU}(p(z) - q(z))}$。
现在再倒回来看,如果拒绝后仍然使用 $p(x)$ 进行采样,则要保持分布不变,必须满足
\[q(y)r(y) + p(y)P(\mathrm{reject}) = p(y),\]于是解得
\[r(y) = \frac{p(y)P(\mathrm{accept})}{q(y)},\]加上 $r(y) \le 1$ 的约束有:
\[r(y) = \min \left( 1, \frac{p(y)P(\mathrm{accept})}{q(y)} \right).\]因此除非 $P(\mathrm{accept}) \equiv 1$,否则真实使用的接受概率 $\min \left( 1, \frac{p(y)}{q(y)} \right)$ 与要保证分布不变所需的接受概率 $\min \left( 1, \frac{p(y)P(\mathrm{accept})}{q(y)} \right)$ 不同。由于 $r(z) \le 1$,若要保证
\[P(\mathrm{accept}) = \sum _{z} q(z) r(z) = 1,\]则对于任意使得 $q(z) > 0$ 的 $z$,均必须有 $r(z) = 1$,即 $q(z) = p(z)$。于是易知必须要 $p(x) = q(x)$ for all $x \in V$。综上:如果以概率 $\min \left( 1, \frac{p(y)}{q(y)} \right)$ 接受 token,则除非 $q \equiv p$,否则拒绝后使用 $p$ 采用无法保证分布不偏移。
那么如果我们就是想要在拒绝后按分布 $p$ 采样 token 呢(虽然实操的角度没必要,因为这样做没办法省时间,但是可以讨论一下),那么我们就需要调整接受策略 $r(y)$。记接受概率 $P(\mathrm{accept}) = A(r)$(它与 $r$ 相关),且当 $p$ 不恒等于 $q$ 时,$A(r)<1$。于是我们有
\[r(y) = \min \left( 1, \frac{p(y)A(r)}{q(y)} \right).\]首先,这样做的接受概率就天然小于选择以概率 $\min \left( 1, \frac{p(y)}{q(y)} \right)$ 接受 token 时的接受概率,从而会导致接受长度变短;其次,由于上式右边与 $r$ 有关,要想从上式中解出 $r$ 真正的表达式也比较困难。因此,不值得坚持使用 $p$ 去在拒绝后进行采样。