quant.nx source
↩ module page · 462 lines · 16869 B
1// quant.nx -- quantization primitives for AI inference.
2//
3// Packs / unpacks tensor data between fp16/fp32 and the
4// low-precision formats nxgguf supports (int8, int4, int2,
5// ternary, fp8 E4M3). These are the actual VRAM-saving
6// transforms that take a 32GB Llama 70B fp16 model down to
7// 8GB int4 or 4GB int2.
8//
9// Reference quantization schemes:
10// int8 SmoothQuant (Xiao et al. 2022)
11// int4 GPTQ (Frantar et al. 2022)
12// int4 AWQ (Lin et al. 2023)
13// GGUF k-quants (llama.cpp project)
14// BitNet b1.58 (Wang et al. 2024) -- ternary {-1, 0, +1}
15//
16// v0.0.1 ships symmetric quantization with per-tensor scale.
17// k-quants (per-block scale + zero-point) follow in v0.1.0.
18
19import "syscalls.nx"
20
21// === fp32 literal encoding ===========================================
22//
23// Convert a lexer-split decimal literal (whole + fractional-digits +
24// number-of-fractional-digits) to IEEE 754 binary32 bit pattern.
25// Called from parse.nx when it sees TK_FLOAT and needs an IR
26// constant with a concrete bit pattern.
27//
28// Algorithm:
29// value = (whole * 10^frac_digits + frac_num) / 10^frac_digits
30// scale the numerator by 2^40 before the divide to preserve
31// precision, locate the highest set bit to derive the exponent,
32// then pack mantissa+exponent into the 32-bit layout.
33//
34// Precision: accurate to ~7 decimal digits (f32 is 24-bit
35// significand). Exact for common cases like 0.0, 0.5, 1.0, 2.0,
36// 1.5, 3.14, 0.1 (all tested). Overflow saturates to +inf;
37// underflow flushes to zero.
38//
39// Negative literals handled by parse.nx prepending a unary minus
40// separately -- the lexer emits only the unsigned magnitude.
41
42func fp32_from_parts(whole: i64, frac_num: i64, frac_digits: i64) -> i64 {
43 // Zero literal -> bit pattern 0.
44 if whole == 0 {
45 if frac_num == 0 { return 0 }
46 }
47
48 // 10^frac_digits.
49 var denom: i64 = 1
50 var i: i64 = 0
51 while i < frac_digits {
52 denom = denom * 10
53 i = i + 1
54 }
55
56 // Scaled numerator. 2^40 gives room for a 24-bit mantissa above
57 // denom across the common literal range -- 40-bit of margin
58 // drowns rounding error for 7-digit literals.
59 let num: i64 = whole * denom + frac_num
60 let scale: i64 = 40
61 let scaled: i64 = num << scale
62 let quot: i64 = scaled / denom
63
64 // Locate highest set bit of quot.
65 var high_bit: i64 = 0
66 var q: i64 = quot
67 while q > 1 {
68 q = q >> 1
69 high_bit = high_bit + 1
70 }
71
72 // True exponent = high_bit - scale.
73 let true_exp: i64 = high_bit - scale
74
75 // Align quot so the implicit 1 bit lands at bit 23, then mask.
76 var mant: i64 = 0
77 if high_bit >= 23 {
78 mant = quot >> (high_bit - 23)
79 }
80 if high_bit < 23 {
81 mant = quot << (23 - high_bit)
82 }
83 mant = mant & 0x7FFFFF
84
85 // Bias exponent by 127. Saturate on overflow, flush on underflow.
86 let biased: i64 = true_exp + 127
87 if biased < 0 { return 0 }
88 if biased > 254 { return 0x7F800000 }
89 return (biased << 23) | mant
90}
91
92// === fp64 literal encoding ===========================================
93//
94// Convert a lexer-split decimal literal to IEEE 754 binary64 bit
95// pattern. Mirrors fp32_from_parts; differs in:
96// - mantissa width 52 (vs 23)
97// - exponent bias 1023 (vs 127)
98// - exponent saturation at 2046 / +inf bit pattern 0x7FF0000000000000
99// - scale 52 (vs 40) -- gives 52-bit mantissa headroom in the
100// scaled-numerator divide. Caps decodable literal magnitude at
101// ~3 digits whole + ~3 digits frac before the i64 multiply
102// overflows; matches fp32_from_parts's same pragmatic limit.
103// Larger-magnitude literals need i128 multiword arithmetic
104// (queued: u128/i128 substrate primitives, see roadmap doc 16).
105//
106// Precision: accurate to ~7 decimal digits today (limited by the
107// scale=52 i64 budget, not by the algorithm itself). Exact for
108// 0.0, 0.5, 1.0, 2.0, 1.5, 4.0, 8.0; ULP-close for 0.1, 3.14.
109//
110// Negative literals: parse.nx prepends unary minus separately --
111// the lexer emits unsigned magnitude only.
112//
113// IEEE 754 binary64 layout: 1 sign + 11 exp + 52 mant.
114
115func fp64_from_parts(whole: i64, frac_num: i64, frac_digits: i64) -> i64 {
116 // Zero literal -> bit pattern 0.
117 if whole == 0 {
118 if frac_num == 0 { return 0 }
119 }
120
121 // 10^frac_digits.
122 var denom: i64 = 1
123 var i: i64 = 0
124 while i < frac_digits {
125 denom = denom * 10
126 i = i + 1
127 }
128
129 // Scaled numerator. 2^52 of margin gives 52-bit mantissa above
130 // denom for typical literal magnitudes.
131 let num: i64 = whole * denom + frac_num
132 let scale: i64 = 52
133 let scaled: i64 = num << scale
134 let quot: i64 = scaled / denom
135
136 // Locate highest set bit of quot.
137 var high_bit: i64 = 0
138 var q: i64 = quot
139 while q > 1 {
140 q = q >> 1
141 high_bit = high_bit + 1
142 }
143
144 // True exponent = high_bit - scale.
145 let true_exp: i64 = high_bit - scale
146
147 // Align quot so the implicit 1 bit lands at bit 52, then mask off
148 // the implicit bit (IEEE stores only the trailing 52 bits).
149 var mant: i64 = 0
150 if high_bit >= 52 {
151 mant = quot >> (high_bit - 52)
152 }
153 if high_bit < 52 {
154 mant = quot << (52 - high_bit)
155 }
156 mant = mant & 0xFFFFFFFFFFFFF
157
158 // Bias exponent by 1023. Saturate on overflow, flush on underflow.
159 let biased: i64 = true_exp + 1023
160 if biased < 0 { return 0 }
161 if biased > 2046 { return 0x7FF0000000000000 }
162 return (biased << 52) | mant
163}
164
165// === fp64 -> fp32 downcast ==========================================
166//
167// Round-to-nearest-even IEEE 754 binary64 -> binary32 conversion.
168// Used by parse_stmt_let when type-context inference re-types a
169// default-f64 literal at a `let x: f32 = ...` site. Round-trip
170// stable for f32-representable values (e.g. 1.5, 2.0, 0.5); rounds
171// to nearest f32 representation for others.
172//
173// Underflow (true exponent < -126) flushes to zero; overflow
174// (true exponent > 127) saturates to +/- inf. NaN bit pattern
175// transfers through with a single bit set in the f32 mantissa.
176
177func fp64_to_fp32(bits64: i64) -> i64 {
178 let sign: i64 = (bits64 >> 32) & 0x80000000
179 let exp: i64 = (bits64 >> 52) & 0x7FF
180 let mant: i64 = bits64 & 0xFFFFFFFFFFFFF
181 if exp == 0 {
182 // Zero or subnormal f64 -- f32 underflows; flush to zero.
183 return sign
184 }
185 if exp == 0x7FF {
186 // Inf or NaN.
187 if mant == 0 { return sign | 0x7F800000 }
188 return sign | 0x7F800000 | 1
189 }
190 let true_exp: i64 = exp - 1023
191 let new_exp: i64 = true_exp + 127
192 if new_exp <= 0 {
193 return sign
194 }
195 if new_exp >= 0xFF {
196 return sign | 0x7F800000
197 }
198 // f64 mantissa is 52 bits; f32 wants 23. Drop the low 29 bits
199 // with round-to-nearest-even.
200 let drop: i64 = 29
201 let lo: i64 = mant & ((1 << drop) - 1)
202 let half: i64 = 1 << (drop - 1)
203 var new_mant: i64 = mant >> drop
204 // Round-to-nearest-even tie-break: if dropped bits == half AND
205 // new_mant is even, round down; else round up.
206 if lo > half {
207 new_mant = new_mant + 1
208 }
209 if lo == half {
210 if (new_mant & 1) == 1 { new_mant = new_mant + 1 }
211 }
212 // Mantissa overflow into exponent (e.g., 1.111... rounds up to 10.0).
213 var final_exp: i64 = new_exp
214 if new_mant >= (1 << 23) {
215 new_mant = new_mant >> 1
216 final_exp = final_exp + 1
217 if final_exp >= 0xFF { return sign | 0x7F800000 }
218 }
219 new_mant = new_mant & 0x7FFFFF
220 return sign | (final_exp << 23) | new_mant
221}
222
223// === fp16 / bf16 conversion ==========================================
224//
225// IEEE 754 binary16 layout: 1 sign + 5 exponent + 10 mantissa.
226// Bias = 15. Subnormals + infinity + NaN encoded standardly.
227
228// Convert IEEE 754 fp32 (passed as i64 holding the bit pattern)
229// to fp16 (returned as i64 holding 16 bits). Round-to-nearest-even.
230// Underflow flushes to zero; overflow saturates to +/- inf.
231func fp32_to_fp16(bits32: i64) -> i64 {
232 let sign: i64 = (bits32 >> 16) & 0x8000
233 let mant: i64 = bits32 & 0x7FFFFF
234 let exp: i64 = (bits32 >> 23) & 0xFF
235 if exp == 0 { return sign } // zero or subnormal -> 0
236 if exp == 0xFF { // inf or NaN
237 if mant == 0 { return sign | 0x7C00 } // inf
238 return sign | 0x7C00 | 1 // NaN (encode any non-zero mant)
239 }
240 let new_exp: i64 = exp - 127 + 15 // bias adjust
241 if new_exp <= 0 { // too small -> flush to zero
242 return sign
243 }
244 if new_exp >= 0x1F { // too large -> inf
245 return sign | 0x7C00
246 }
247 let new_mant: i64 = mant >> 13
248 return sign | (new_exp << 10) | new_mant
249}
250
251// fp16 -> fp32 inverse.
252func fp16_to_fp32(bits16: i64) -> i64 {
253 let sign: i64 = (bits16 & 0x8000) << 16
254 let exp: i64 = (bits16 >> 10) & 0x1F
255 let mant: i64 = bits16 & 0x3FF
256 if exp == 0 {
257 if mant == 0 { return sign } // zero
258 // Subnormal fp16 -- normalize to fp32.
259 var e: i64 = 1
260 var m: i64 = mant
261 var top_bit: i64 = m & 0x400
262 while top_bit == 0 {
263 m = m << 1
264 e = e + 1
265 top_bit = m & 0x400
266 }
267 let new_exp: i64 = 127 - 15 - e + 1
268 return sign | (new_exp << 23) | ((m & 0x3FF) << 13)
269 }
270 if exp == 0x1F {
271 if mant == 0 { return sign | 0x7F800000 } // inf
272 return sign | 0x7F800000 | (mant << 13) // NaN
273 }
274 let new_exp: i64 = exp - 15 + 127
275 return sign | (new_exp << 23) | (mant << 13)
276}
277
278// === fp8 E4M3 conversion =============================================
279//
280// 1 sign + 4 exponent + 3 mantissa. Bias = 7. Used by H100 + RTX 50
281// for inference. NaN encoding: all 1s in exp + non-zero mant.
282// Doesn't have inf (saturates).
283
284func fp32_to_fp8e4m3(bits32: i64) -> i64 {
285 let sign: i64 = (bits32 >> 24) & 0x80
286 let mant: i64 = bits32 & 0x7FFFFF
287 let exp: i64 = (bits32 >> 23) & 0xFF
288 if exp == 0 { return sign }
289 if exp == 0xFF {
290 // Inf/NaN -> NaN in fp8 (saturate-to-inf isn't supported in E4M3).
291 return sign | 0x7F
292 }
293 let new_exp: i64 = exp - 127 + 7
294 if new_exp <= 0 { return sign } // flush to zero
295 if new_exp >= 0xF { // saturate to max-finite
296 return sign | (0xE << 3) | 0x7
297 }
298 let new_mant: i64 = mant >> 20
299 return sign | (new_exp << 3) | (new_mant & 0x7)
300}
301
302// === int8 symmetric quantization =====================================
303//
304// Maps fp values in [-max_abs, +max_abs] to int8 [-127, 127].
305// scale = max_abs / 127. Dequant: fp = i8 * scale.
306
307// Find max absolute value across n fp32 elements (passed as bit
308// patterns in an i64 array).
309func find_max_abs_fp32(data: *u8, n: i64) -> i64 {
310 var max_bits: i64 = 0
311 var i: i64 = 0
312 while i < n {
313 let off: i64 = i * 4
314 let bits: i64 =
315 data[off]
316 | (data[off + 1] << 8)
317 | (data[off + 2] << 16)
318 | (data[off + 3] << 24)
319 let abs_bits: i64 = bits & 0x7FFFFFFF
320 if abs_bits > max_bits { max_bits = abs_bits }
321 i = i + 1
322 }
323 return max_bits
324}
325
326// Quantize one fp32 value (bits) to int8 given a precomputed scale
327// (also fp32 bits). Returns int8 stored in i64 [-127, 127].
328// v0.0.1 uses an integer-arithmetic divide approximated by shift +
329// magnitude check; full IEEE float divide lands when the F extension
330// codegen is wired (see REGALLOC_ROADMAP.md).
331func quant_fp32_to_int8(value_bits: i64, scale_bits: i64) -> i64 {
332 // Approximate: extract magnitude, scale by ratio of mantissas.
333 // For v0.0.1, we ship the API + a placeholder that returns the
334 // raw mantissa shifted; real IEEE divide arrives with F-ext.
335 let mant: i64 = value_bits & 0x7FFFFF
336 let sign: i64 = (value_bits >> 31) & 1
337 let smant: i64 = scale_bits & 0x7FFFFF
338 if smant == 0 { return 0 }
339 var q: i64 = (mant * 127) / (smant + 1)
340 if q > 127 { q = 127 }
341 if sign == 1 { q = 0 - q }
342 return q
343}
344
345// Dequantize int8 back to fp32 bits = i8 * scale.
346// Same v0.0.1 caveat: integer-arithmetic placeholder until F-ext.
347func dequant_int8_to_fp32(q: i64, scale_bits: i64) -> i64 {
348 if q == 0 { return 0 }
349 let abs_q: i64 = q
350 var mag: i64 = abs_q
351 if q < 0 { mag = 0 - q }
352 let smant: i64 = scale_bits & 0x7FFFFF
353 let new_mant: i64 = (mag * smant) / 127
354 let exp_part: i64 = scale_bits & 0xFF800000
355 if q < 0 { return 0x80000000 | exp_part | (new_mant & 0x7FFFFF) }
356 return exp_part | (new_mant & 0x7FFFFF)
357}
358
359// === int4 packing ====================================================
360//
361// Two 4-bit nibbles per byte. Low nibble = element 0, high nibble
362// = element 1. Range [-8, 7] symmetric.
363
364// Pack two int4 values (each in i64 holding [-8, 7]) into one byte.
365func pack_int4_pair(lo: i64, hi: i64) -> i64 {
366 let lo_n: i64 = lo & 0xF
367 let hi_n: i64 = hi & 0xF
368 return lo_n | (hi_n << 4)
369}
370
371// Extract the low int4 from a packed byte; sign-extend.
372func unpack_int4_lo(byte: i64) -> i64 {
373 let n: i64 = byte & 0xF
374 if n >= 8 { return n - 16 }
375 return n
376}
377
378// Extract the high int4 from a packed byte; sign-extend.
379func unpack_int4_hi(byte: i64) -> i64 {
380 let n: i64 = (byte >> 4) & 0xF
381 if n >= 8 { return n - 16 }
382 return n
383}
384
385// Pack n int4 values (passed as i8-in-i64 array) into floor(n/2)
386// bytes. Last odd element (if any) is dropped.
387func pack_int4_array(in_vals: *u8, n: i64, out: *u8) -> i64 {
388 var i: i64 = 0
389 var oi: i64 = 0
390 while i + 1 < n {
391 out[oi] = pack_int4_pair(in_vals[i], in_vals[i + 1])
392 i = i + 2
393 oi = oi + 1
394 }
395 return oi
396}
397
398// === int2 / ternary packing ==========================================
399//
400// 4 elements per byte (2 bits each). Ternary (BitNet b1.58) maps
401// {-1, 0, +1} to {0b00, 0b01, 0b10}; 0b11 reserved.
402
403func pack_int2_quad(a: i64, b: i64, c: i64, d: i64) -> i64 {
404 return (a & 0x3) | ((b & 0x3) << 2) | ((c & 0x3) << 4) | ((d & 0x3) << 6)
405}
406
407func unpack_int2_at(byte: i64, idx: i64) -> i64 {
408 return (byte >> (idx * 2)) & 0x3
409}
410
411// Ternary mapping: {-1, 0, +1} -> {0, 1, 2} -> packed 2-bit
412func ternary_encode(v: i64) -> i64 {
413 if v < 0 { return 0 }
414 if v == 0 { return 1 }
415 return 2
416}
417func ternary_decode(b: i64) -> i64 {
418 if b == 0 { return -1 }
419 if b == 1 { return 0 }
420 return 1
421}
422
423// === self-test ===
424
425func main() -> i64 {
426 // fp16 round-trip: 1.0 in fp32 = 0x3F800000.
427 // fp16(1.0) = 0x3C00 (sign 0, exp 15, mant 0).
428 let fp16_one: i64 = fp32_to_fp16(0x3F800000)
429 if fp16_one != 0x3C00 { return __syscall(93, 50, 0, 0, 0, 0, 0) }
430
431 // fp16 -> fp32 round-trip on 1.0.
432 let back: i64 = fp16_to_fp32(0x3C00)
433 if back != 0x3F800000 { return __syscall(93, 51, 0, 0, 0, 0, 0) }
434
435 // fp16 of zero stays zero.
436 if fp32_to_fp16(0) != 0 { return __syscall(93, 52, 0, 0, 0, 0, 0) }
437 // fp16 of -0.0 (sign bit set, rest zero) stays -0.0.
438 if fp32_to_fp16(0x80000000) != 0x8000 { return __syscall(93, 53, 0, 0, 0, 0, 0) }
439
440 // int4 pack/unpack: pack(3, -2) = 0x3 | (0xE << 4) = 0xE3 (-2 = 0xE in 4-bit).
441 let packed: i64 = pack_int4_pair(3, -2)
442 if packed != 0xE3 { return __syscall(93, 60, 0, 0, 0, 0, 0) }
443 if unpack_int4_lo(packed) != 3 { return __syscall(93, 61, 0, 0, 0, 0, 0) }
444 if unpack_int4_hi(packed) != -2 { return __syscall(93, 62, 0, 0, 0, 0, 0) }
445
446 // int2 quad: (1, 2, 3, 0) packs to 0x39 = 0b00111001
447 // = (3<<6) | (3<<4) | (2<<2) | 1 -- wait that's 1 | 2<<2 | 3<<4 | 0<<6
448 // = 0x01 | 0x08 | 0x30 | 0x00 = 0x39
449 let q: i64 = pack_int2_quad(1, 2, 3, 0)
450 if q != 0x39 { return __syscall(93, 70, 0, 0, 0, 0, 0) }
451 if unpack_int2_at(q, 0) != 1 { return __syscall(93, 71, 0, 0, 0, 0, 0) }
452 if unpack_int2_at(q, 1) != 2 { return __syscall(93, 72, 0, 0, 0, 0, 0) }
453 if unpack_int2_at(q, 2) != 3 { return __syscall(93, 73, 0, 0, 0, 0, 0) }
454 if unpack_int2_at(q, 3) != 0 { return __syscall(93, 74, 0, 0, 0, 0, 0) }
455
456 // Ternary: -1 -> 0, 0 -> 1, +1 -> 2 round-trip.
457 if ternary_decode(ternary_encode(-1)) != -1 { return __syscall(93, 80, 0, 0, 0, 0, 0) }
458 if ternary_decode(ternary_encode(0)) != 0 { return __syscall(93, 81, 0, 0, 0, 0, 0) }
459 if ternary_decode(ternary_encode(1)) != 1 { return __syscall(93, 82, 0, 0, 0, 0, 0) }
460
461 return __syscall(93, 42, 0, 0, 0, 0, 0)
462}