code wiki / (root) / nx_q8_0_from_f32_kat.nx

nx_q8_0_from_f32_kat.nx source

↩ module page · 37 lines · 1502 B

1// nx_q8_0_from_f32_kat.nx -- roundtrip KAT: F32 -> Q8_0 -> F32 within a quant 2// step. Catches f32->f16 / scale / rounding bugs. expect_exit: 0. 3import "nx_syscalls.nx" 4import "nx_f32.nx" 5import "nx_f32_cvt.nx" 6import "nx_q8_0_from_f32.nx" 7import "nx_q8_0_to_f32.nx" 8 9func main() -> i64 { 10 let n: i64 = 96 11 let f32in: *i64 = sys_mmap(n * 8) as *i64 12 var i: i64 = 0 13 while i < n { 14 // varied magnitudes: (i-48) scaled by /7 -> non-trivial fractions 15 f32in[i] = nx_f32_div(nx_i32_to_f32(i - 48), nx_i32_to_f32(7)) 16 i = i + 1 17 } 18 let nblk: i64 = (n + 31) / 32 19 let q8: *u8 = sys_mmap(nblk * 34) 20 nx_q8_0_from_f32(f32in, n, q8) 21 let deq: *i64 = sys_mmap(n * 8) as *i64 22 nx_q8_0_to_f32(q8, 0, n, deq) 23 // |deq - orig| must be < 0.5 (each block's d = absmax/127 <= ~0.9 here, and 24 // error <= d/2; a slack bound of 0.5 catches gross scale/encode bugs). 25 i = 0 26 while i < n { 27 // deq - orig = deq + (-orig); negate orig by flipping its sign bit. 28 let diff: i64 = __f32_add(deq[i], f32in[i] ^ 0x80000000) & 0x7FFFFFFF 29 if diff > 0x3F000000 { return 40 + i } // |diff| >= 0.5 -> FAIL 30 i = i + 1 31 } 32 // also: block-0 exact endpoint -- orig[0] = -48/7, absmax in block0 = 48/7, 33 // so int8[0] = -127, dequant = -absmax = -48/7 exactly. verify tight. 34 let d0: i64 = __f32_add(deq[0], f32in[0] ^ 0x80000000) & 0x7FFFFFFF 35 if d0 > 0x3B800000 { return 30 } // > ~0.004 -> endpoint not tight 36 return 0 37}