code wiki / (root) / nx_quant.nx

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}