What disappears when the weights are quartered
Part two ended with the cache larger than the weights. This part goes the other way and shrinks the weights.
The method is simple. A float32 is four bytes, but knowing the range of the values lets one integer carry each of them.
qmax = 2 ** (bits - 1) - 1
scale = w.abs().max() / qmax
q = (w / scale).round().clamp(-qmax - 1, qmax) # the integers
w_hat = q * scale # unpacked again
Keep one float scale and hold the rest as int8, and four bytes become one. The
question is how much the round hurts.
int8 is nearly free
Using part thirteen’s model unchanged, its float32 validation loss is 2.0272.
one scale per matrix one scale per channel
8 bits 2.0290 2.0280
float32 is 2.0272
+0.0017. It barely shows in the third decimal. All 637k weights were flattened
onto 256 levels and the model hardly noticed.
The memory goes like this.
631,808 two-dimensional weights 2.41 MB (fp32) -> 0.60 MB (int8)
all 637,156 2.43 MB -> 0.62 MB 74.4% less
Squeeze further and where does it break
bits one scale per matrix one scale per channel
8 2.0290 2.0280
6 2.0599 2.0387
4 2.6380 2.3571
3 4.3960 3.4381
2 5.7894 6.3112
Through 6 bits it is still a second-decimal conversation. At 4 it separates to
2.36 and at 3 it collapses to 3.44. Part one’s bigram scored 2.6501, so
a four-bit model is barely better than counting the previous character.
At two bits it inverts
The last row is odd. At 2 bits the per-channel scale gives 6.3112, which is
worse than the per-matrix 5.7894 - after being better on every row above.
Measuring what happened, for blocks.0.f1 alone:
one scale per matrix 95.2% of elements collapse to zero mean abs error 0.0519
one scale per channel 78.1% collapse to zero mean abs error 0.0419
Per-channel reconstructed the weights better. Fewer dead elements, smaller error. And the loss is worse.
So at two bits, “how well were the weights reconstructed” and “how well does the model work” have already come apart. Neither is usable, and the ordering between them means nothing. The right move is to write down what was measured and decline to draw a conclusion from it.
What is sensitive
Instead of squeezing everything, take one matrix down to 4 bits and leave the rest at float32, and it becomes clear who hurts.
tok.weight loss +0.0976
pos.weight loss +0.0618
blocks.2.f1.weight loss +0.0184
blocks.0.f2.weight loss +0.0174
blocks.2.qkv.weight loss +0.0132
...
blocks.2.f2.weight loss -0.0020 (least sensitive; it dips slightly)
The two embeddings dominate. The character embedding alone hurts more than five times as much as any matrix inside a block. This is why real quantisation keeps embeddings and the output layer at higher precision.
Worth noting that the least sensitive is -0.0020, a negative number.
Quantisation did not accidentally help; a change that size is inside the wobble
of the loss measurement itself.
Adding up the individual damages gives +0.2830 while quantising everything
gives +0.3299. The damage does not merely add; the pieces amplify each other.
The loss holds still and the text does not
int8 moved the loss by +0.0017. Does the same text come out? Generating from the
same seed:
float32 '`16, a=0, 36, 112003 0.3009849\n```\n\nThe smallest `1.000h`. **`shape(2'
int8 '`16, a=0, 36, 112003 20350 0.63103 20.013910 0.016693939 16\n9 = '
It diverges at the twenty-second character. Identical up to there and a completely different text after.
Which follows. As part one showed, candidates sit close together, and where the
probability gap is tiny a 0.0017 wobble picks a different character. Change one
character and the whole context after it changes.
Part two’s note about the cache’s 9.06e-06 - that it is not bit-identical - here
becomes a visible consequence. Equal average performance and identical output
are different claims.
Per-channel barely mattered in this model
The usual argument in quantisation is about outliers: if one channel’s range is unusually wide, a single scale per matrix flattens everything else. Hence splitting by channel.
Measured in this model, that premise is weak.
ratio of max to median channel range worst matrix 2.49 median 1.51
About 1.5. At that spread a single scale costs little, and indeed the two
schemes differ by 0.0010 at 8 bits.
The activations say the same.
block 0 FFN input channel max/median 1.50 max 5.24
block 2 FFN input channel max/median 1.54 max 6.24
residual stream (last) channel max/median 1.64 max 6.67
What makes activation quantisation hard in large models is a handful of channels
spiking to tens or hundreds of times the rest; here it is 1.6. So this
experiment does not reproduce that problem. The fair reading is that part
thirteen’s model is small and was trained for 5000 steps on sixty thousand
characters, which means this part’s finding that “per-channel adds little” is
specific to this model.
So
- int8 is nearly free: validation loss
2.0272->2.0290, memory2.43 MB->0.62 MB,74.4%less - It holds through
6bits. At4it reads2.3571, barely past the bigram’s2.6501; at3it collapses - At
2bits the scheme that reconstructed the weights better has the worse loss. The two measures have already parted and the ordering means nothing - Embeddings are the most sensitive. At 4 bits
tokcosts+0.0976, over five times any matrix inside a block - A
+0.0017loss still diverges the text at the twenty-second character. Equal averages are not identical outputs - This model has no outliers. A channel-range ratio of
1.5leaves little for per-channel to win, so this conclusion does not transfer to a large model
Each time what the theory promised and what arrived were different.
Next time the 4-bit model built here gets reused rather than discarded. A cheap model writes characters ahead and an expensive one checks them in a single pass - starting with whether the output really survives that.
Comments