The better the draft, the slower it got
In part two, every character required one pass through the target model. Adding a cache does nothing about that: the number of passes still equals the number of characters.
Speculative decoding attacks the count. A cheap model writes k characters
ahead, and the expensive model checks all k+1 positions in one pass.
Whatever matched is kept and everything from the first mismatch is thrown away.
It comes with a claim that sounds too good: the output distribution is exactly
what the target alone would have produced. Given that part three’s int8 moved
the loss by 0.0017 and still diverged the text at the twenty-second character,
that deserves checking first.
Why the distribution survives
Say the draft proposed x with draft probability q(x) against target
probability p(x). The whole rule is this.
if rand() < min(1, p[x] / q[x]):
accept(x)
else:
resid = (p - q).clamp_min(0)
emit(sample(resid / resid.sum())) # and discard the rest
break
Where q proposed x more often than p wanted it, only that fraction is kept;
where it is rejected, the replacement is drawn from max(0, p-q), the part p
wanted and q under-supplied. The two cases together come to exactly p.
Measured rather than argued. Fixing a context and drawing the first character 20000 times, against the target distribution:
bigram draft total variation 0.0158
4-bit draft total variation 0.0095
sampling the target total variation 0.0134 <- baseline at the same count
Right at the baseline. Twenty thousand draws carry that much sampling error anyway, and what speculative decoding adds sits inside it. However bad the draft model is, the output belongs to the target.
Three drafts
Three drafts, all of them already in this series.
- bigram: part one’s baseline, a table lookup on the previous character
- 8-bit: part three’s weights flattened to int8
- 4-bit: part three’s, one step from collapse
Acceptance follows draft quality
draft k acceptance chars per target pass target passes
bigram 1 0.290 1.29 155
bigram 8 0.059 1.47 136
8-bit 1 0.896 1.89 106
8-bit 8 0.610 5.88 34
4-bit 1 0.653 1.65 121
4-bit 8 0.351 3.77 53
As expected. The 8-bit draft is nearly the target itself and gets 0.896
accepted; the bigram gets 0.290. Raising k lowers acceptance - the context
drifts further out with each drafted character - while raising the characters
harvested per pass.
At k=8 the 8-bit draft takes the target from 200 passes down to 34, 5.88
characters each. If part two’s cache changed the exponent, this divides the
constant by six.
And in practice the 8-bit draft is the slowest
draft k chars per pass measured range
bigram 4 1.39 1.30 1.19-1.47
bigram 8 1.47 1.26 1.08-1.36
8-bit 4 4.35 0.80 0.79-0.91
8-bit 8 5.88 0.56 0.50-0.66
4-bit 8 3.77 0.41 0.38-0.47
The one that cut target passes 5.9-fold came out 0.56 times as fast. The
bigram, which cut them by barely 1.47, runs 1.26 times faster.
The reason is what the draft costs. Timing one drafted character against one target pass:
one target pass 1235 us
one bigram character 5 us c = 0.004
one 8-bit character 1307 us c = 1.059
one 4-bit character 1276 us c = 1.034
c = 1.059. The 8-bit draft is more expensive than the target. So at k=8,
saving one target pass costs eight draft passes, and those eight cost what eight
target passes would.
Re-reading part three makes this obvious. What int8 saved there was memory:
2.43 MB became 0.62 MB, and nothing at all was claimed about time. The
multiplications still run in float32 with the weights merely rounded. Reading
quantisation as a speed win produces exactly this mistake.
A formula that holds
For the first time in this series, a simple expression matched the measurement.
speedup = (chars per target pass) / (1 + k · c)
What one target pass yields on top, what was paid in draft calls to get it underneath.
draft k predicted measured median
bigram 2 1.36 1.10
bigram 4 1.37 1.30
8-bit 2 0.87 0.87
8-bit 4 0.83 0.80
8-bit 8 0.62 0.56
4-bit 2 0.67 0.67
4-bit 4 0.49 0.46
4-bit 8 0.41 0.41
For the quantised drafts it holds to the second decimal. For the bigram it runs a
little generous, because c was timed from a single call and leaves out the
Python loop around it - and with a cost that small, that share looks large.
Set this against part two, where the arithmetic promised 134.5 and the
measurement gave 2.94. The difference is what went into the count: there,
only multiplications; here, everything that was paid.
A note on the timings
The timings in this part were untrustworthy at first. The same baseline read
749 ms on one run and 448 ms on the next, because the laptop’s load moves.
So the baseline and the speculative run were measured alternately, a ratio
taken per pair, and the median of five pairs reported. Drift hits both sides of a
ratio and cancels. The range column is the min and max of those five, and some,
like the bigram at k=1, are as wide as 0.82-1.29. The band in the figure is
that width.
Acceptance rates and pass counts are deterministic given the seed and have none of this problem, which is why this part’s conclusion leans on those rather than on the clock.
So
- Speculative decoding does not change the output distribution: total variation
0.0158and0.0095against a same-sample-count baseline of0.0134 - A better draft raises acceptance. 8-bit reaches
0.896, the bigram0.290 - Target passes drop a lot. At
k=8the 8-bit draft takes 200 down to 34,5.88characters per pass - And that one is the slowest, at
0.56. The draft costs more than the target, withcat1.059 - The only draft that actually wins is the worst one: the bigram at
1.26, withc = 0.004 - What sets the sign is not acceptance but the draft’s cost:
(chars per pass) / (1 + k·c)
Next time returns to the gap part two could not explain. If overhead is what is being paid for, how large is that share exactly - and is there a way to take it back.
Comments