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