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
We have a slow, big target modelMp, inference from which we're trying to accelerate (distribution p(x)). We also have a quick, efficient approximation modelMq (distribution q(x)). Core idea:
We use the efficient model Mq to generate γ completions.
We use the target model Mp to evaluate all these guesses inparallel, accepting all those that can lead to an identical distribution.
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), we accept this sample
Or, we reject it with probability 1−q(x)p(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 tokens in a parallel run of the target model Mp (the worst is 1 token, but we know the traditional method also get only 1 token).
This operation is being done in parallel. We concat prefix with the γ guess tokens to make a complete array and input to the target model Mp all at once.
Causal Mask mechanism of the Transformer allows model to output the probability distribution of these γ+1 positions in parallel during one signal forward pass computation.
r1∼U(0,1),...,rγ∼U(0,1)
n←min({i−1∣1≤i≤γ,ri>qi(x)pi(x)}∪{γ})
Generate γ uniformly distributed random numbers between 0and 1, then we check these numbers to see if token xi need to be rejected.
If pi(xi)≥qi(xi), we can see qi(x)pi(x) will be greater than 1, so it won't be rejected.
n 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)))
If n<γ, meaning the (n+1)the guess was rejected. To correct this error, we do the above operation to generate a new distribution p′(x).
Definition: c is the ratio between the time for a single run of Mq and the time for a single run of Mp
The expected improvement factor in total walltime is (1−α)(γc+1)1−αγ+1
Proof.
Let's say T is the cost of running a single step of Mp. For each run of Algorithm 1 costs Tcγ+T (running Mqγ times and running Mp once). On average we produce 1−α1−αγ+1 tokens. So the overall expected cost for producing a token with Algorithm 1 is tmp=(Tcγ+T)/1−α1−αγ+1=1−αγ+1(cγ+1)(1−α)T. So the factor is T/tmp=(1−α)(γc+1)1−αγ+1
Number of Arithmetic Operations
Similar to walltime improvement so I just pass this one.
Choosing γ
The optimal γ should be the one maximizing the walltime improvement equation (enough compute resources).
[!NOTE]
This γ can be modified dynamically according to the value of β
Appendix
To understand this proof, you need first read the Corollary in Calculate α about β.
First we have p′(x)=norm(max(0,p(x)−q(x)))=∑x′(p(x′)−min(q(x′),p(x′)))p(x)−min(q(x),p(x))=1−βp(x)−min(q(x),p(x))