code wiki / (root) / quant.nx

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}