Quantization
AI & ML: lesson 18 of 32
Store the weights in fewer bits and buy back memory bandwidth.
Lesson 18 of 32 · 5 min
Quantization
Step 1 of 9
One weight, stored in sixteen bits. Every layer holds millions of these and they all have to be read per token.
The Idea
A weight stored in sixteen bits can be stored in eight, or four, by keeping one scale factor per group and rounding each value to the nearest step. Decoding is limited by how fast weights can be read from memory, so fewer bytes per weight is directly fewer milliseconds per token.
Real-World Example
A recipe that says "a pinch" rather than 0.834 grams. The dish still works, the book is thinner, and the cook is faster — right up to the point where the pinch is the leavening.
The Code
w = [-0.82, -0.11, 0.0, 0.37, 0.95]
scale = max(abs(x) for x in w) / 127 # one scale for the group
q = [round(x / scale) for x in w]
print(q) # [-110, -15, 0, 49, 127]
back = [v * scale for v in q]
print(round(max(abs(a - b) for a, b in zip(w, back)), 4)) # 0.0035
q4 = [round(x / (max(abs(x) for x in w) / 7)) for x in w]
print(q4) # [-6, -1, 0, 3, 7] -- 4-bit, coarserYour turn
Fill in the blank.
w = [0.0, 0.5, 1.0]
scale = max(w) / ___
print([round(x / scale) for x in w]) # [0, 64, 127]Mini quiz
1 / 3