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}