NORTHLINE
← All recordsRECORD / 046
Speculative Decoding

What is Sepeculative Decoding

We have a slow, big target model M_p, inference from which we're trying to accelerate (distribution p(x)). We also have a quick, efficient approximation model M_q (distribution q(x

Speculative Decoding14 min read

Ref:Fast Inference from Transformers via Speculative Decoding

Speculative Decoding

Overview

We have a slow, big target model MpM_p, inference from which we're trying to accelerate (distribution p(x)p(x)). We also have a quick, efficient approximation model MqM_q (distribution q(x)q(x)). Core idea:

  1. We use the efficient model MqM_q to generate γ\gamma completions.
  2. We use the target model MpM_p to evaluate all these guesses in parallelin\space parallel, accepting all those that can lead to an identical distribution.
  3. If one token is rejected, we make adjustments to the distribution and then sampling an additional token.

Speculative Sampling

For each sample token:

  • If q(x)≤p(x)q(x)\leq p(x), we accept this sample
  • Or, we reject it with probability 1−p(x)q(x)1-\frac{p(x)}{q(x)}.
    • If it's rejected, we also reject all tokens after this one and make an adjustment to the distribution and sample one correct token.

In the best-case scenario, we can generate γ+1\gamma+1 tokens in a parallel run of the target model MpM_p (the worst is 11 token, but we know the traditional method also get only 11 token).

Pseudo-code:

image-20260317193727526

Here are something notable:

  • p1(x),...,pγ+1←Mp(prefix),...,Mp(prefix+[x1,...,xγ])p_1(x),...,p_{\gamma + 1}\leftarrow M_p(prefix),...,M_p(prefix+[x_1,...,x_\gamma])

    This operation is being done in parallel. We concat prefixprefix with the γ\gamma guess tokens to make a complete array and input to the target model MpM_p all at once.

    Causal Mask mechanism of the Transformer allows model to output the probability distribution of these γ+1\gamma+1 positions in parallel during one signal forward pass computation.

  • r1∼U(0,1),...,rγ∼U(0,1)r_1\sim U(0,1),...,r_\gamma \sim U(0,1)

    n←min({i−1∣1≤i≤γ,ri>pi(x)qi(x)}∪{γ})n\leftarrow \text{min}(\{i-1\mid1\leq i\leq \gamma,r_i\gt\frac{p_i(x)}{q_i(x)}\}\cup\{\gamma\})

    Generate γ\gamma uniformly distributed random numbers between 00and 11, then we check these numbers to see if token xix_i need to be rejected.

    If pi(xi)≥qi(xi)p_i(x_i)\geq q_i(x_i), we can see pi(x)qi(x)\frac{p_i(x)}{q_i(x)} will be greater than 11, so it won't be rejected.

    nn is the index immediately preceding the first rejected position (also it's the number of consecutively accepted tokens).

  • p′(x)←norm(max(0,pn+1(x)−qn+1(x)))p'(x)\leftarrow \text{norm}(\text{max}(0,p_{n+1}(x)-q_{n+1}(x)))

    If n<γn\lt\gamma, meaning the (n+1)(n+1)the guess was rejected. To correct this error, we do the above operation to generate a new distribution p′(x)p'(x).

Prove that xx sampled in this way indeed x∼p(x)x\sim p(x)

See Appendix

Analysis

Number of Generated Tokens

To analyze the expected number of tokens produced by a single run of Algorithm 1, we first have the following definition:

  • The acceptance rate βx<tacceptance\space rate \space \beta_{x\lt t}, given a prefix x<tx_{\lt t}, is the probability of accepting xt∼q(xt∣x<t)x_t\sim q(x_t\mid x_{\lt t}) by speculative sampling.
  • α=E(β)\alpha = E(\beta) is to show how well MqM_q approximates MpM_p

The result is actually a geometric variable. Let's say that we accept kk tokens, the probability of this case:

E(#generated tokens)=1−αγ+11−αE(\#generated\space tokens)=\frac{1-\alpha^{\gamma+1}}{1-\alpha}

image-20260318000325380

[!tip]

We get see E(#generated tokens)E(\#generated\ tokens) as the sum of contribution of each position:

  • For the first token, the contribution is 1 because MpM_p always generates at least once token.
  • For the second token, the contribution is α\alpha, meaning the first one is accepted.
  • For the (γ+1)(\gamma+1)th token, the contribution is αγ\alpha^\gamma
  • So: E(#generated tokens)=1+α+α2+...+αγ=1−αγ+11−αE(\#generated\ tokens)=1+\alpha+\alpha^2+...+\alpha^\gamma=\frac{1-\alpha^{\gamma+1}}{1-\alpha}

Calculating α\alpha

  • Definition: DLK(p,q)=∑x∣p(x)−M(x)∣=∑x∣q(x)−M(x)∣D_{LK}(p,q)=\sum_x\mid p(x)-M(x)\mid=\sum_x\mid q(x)-M(x)\mid where M(x)=p(x)+q(x)2M(x)=\frac{p(x)+q(x)}{2}

  • Lemma: DLK(p,q)=1−∑xmin(p(x),q(x))D_{LK}(p,q)=1-\sum_x\text{min}(p(x),q(x))

    ProofProof. DLK(p,q)=∑x∣p(x)−M(x)∣=∑x∣p−q∣2=1−∑xp+q−∣p−q∣2=1−∑xmin(p(x),q(x))D_{LK}(p,q)=\sum_x\mid p(x)-M(x)\mid=\sum_x\frac{\mid p-q\mid}{2}=1-\sum_x\frac{p+q-\mid p-q\mid}{2}=1-\sum_x\text{min}(p(x),q(x))

[!tip]

In case you forget, ∑p=∑q=1\sum p=\sum q = 1, so we have ∑p+q2=∑p2+∑q2=1\sum\frac{p+q}{2}=\sum\frac{p}{2}+\sum\frac{q}{2}=1

  • Corollary:

    • DLK(p,q)=0  ⟺  p=qD_{LK}(p,q)=0\iff p=q
    • DLK(p,q)=1  ⟺  p and q have disjoint supportD_{LK}(p,q)=1\iff p\ and\ q\ have\ disjoint\ support
  • Theorem: β=1−DLK(p,q)\beta=1-D_{LK}(p,q)

    Proof.Proof. β=Ex∼q(x){1q(x)≤p(x)p(x)q(x)q(x)>p(x)=Ex∼q(x)min(1,p(x)q(x))=∑xmin(p(x),q(x))\beta=E_{x\sim q(x)}\begin{cases}1&q(x)\le p(x)\\\frac{p(x)}{q(x)}&q(x)\gt p(x)\end{cases}=E_{x\sim q(x)}\text{min}(1,\frac{p(x)}{q(x)})=\sum_x \text{min}(p(x),q(x))

  • Corollary: α=1−E(DLK(p,q))=E(min(p,q))\alpha=1-E(D_{LK}(p,q))=E(\text{min}(p,q))

Walltime Improvement

  • Definition: cc is the ratio between the time for a single run of MqM_q and the time for a single run of MpM_p

  • The expected improvement factor in total walltime is 1−αγ+1(1−α)(γc+1)\frac{1-\alpha^{\gamma+1}}{(1-\alpha)(\gamma c+1)}

    Proof.Proof.

    Let's say TT is the cost of running a single step of MpM_p. For each run of Algorithm 11 costs Tcγ+TTc\gamma+T (running MqM_q γ\gamma times and running MpM_p once). On average we produce 1−αγ+11−α\frac{1-\alpha^{\gamma+1}}{1-\alpha} tokens. So the overall expected cost for producing a token with Algorithm 11 is tmp=(Tcγ+T)/1−αγ+11−α=(cγ+1)(1−α)1−αγ+1Ttmp=(Tc\gamma+T)/\frac{1-\alpha^{\gamma+1}}{1-\alpha}=\frac{(c\gamma+1)(1-\alpha)}{1-\alpha^{\gamma+1}}T. So the factor is T/tmp=1−αγ+1(1−α)(γc+1)T/tmp=\frac{1-\alpha^{\gamma+1}}{(1-\alpha)(\gamma c+1)}

Number of Arithmetic Operations

Similar to walltime improvement so I just pass this one.

Choosing γ\gamma

The optimal γ\gamma should be the one maximizing the walltime improvement equation (enough compute resources).

[!NOTE]

This γ\gamma can be modified dynamically according to the value of β\beta

image-20260318165133616

Appendix

To understand this proof, you need first read the Corollary in Calculate α\alpha about β\beta.

First we have p′(x)=norm(max(0,p(x)−q(x)))=p(x)−min(q(x),p(x))∑x′(p(x′)−min(q(x′),p(x′)))=p(x)−min(q(x),p(x))1−βp'(x)=norm(max(0,p(x)-q(x)))=\frac{p(x)-min(q(x),p(x))}{\sum_{x'}(p(x')-min(q(x'),p(x')))}=\frac{p(x)-min(q(x),p(x))}{1-\beta}

Now:

P(x=x′)=P(guess accepted,x=x′)+P(guess rejected,x=x′)P(x=x')=P(guess\ accepted,x=x')+P(guess\ rejected,x=x')

Where:

P(guess accepted,x=x′)=q(x′)min(1,p(x′)q(x′))=min(q(x′),p(x′))P(guess\ accepted,x=x')=q(x')min(1,\frac{p(x')}{q(x')})=min(q(x'),p(x'))

And:

P(guess rejected,x=x′)=p′(x′)(1−β)=p(x′)−min(q(x′),p(x′))P(guess\ rejected,x=x')=p'(x')(1-\beta)=p(x')-min(q(x'),p(x'))

Overall:

P(x=x′)=p(x′)P(x=x')=p(x')
End of record / 046
← All records
READ NEXT[Note] FreeKV: Boosting KV Cache Retrieval For Efficient LLM Inference