code wiki / _hdl_build / nx_mul_wide_test.nx
nx_mul_wide_test.nx source
↩ module page · 134 lines · 5083 B
1// nx_mul_wide_test.nx -- TRIANGULATED validation of the 64x64->128 multiplier.
2// Three INDEPENDENT legs must agree (a bug would have to fool all three):
3// LEG A gate-network (32-bit-limb schoolbook) -> (hiA, loA)
4// LEG B independent 128-bit shift-add reference (bit-by-bit, two-word
5// accumulator w/ unsigned carry -- a DIFFERENT algorithm) -> (hiB, loB)
6// LEG C modular Freivalds: hi*2^64+lo == a*b (mod p) over 3 primes
7// + sanity loA == a*b (the i64 product = the true low 64 bits)
8// Triangulation (operator: build bottom-up, cross-check with both 64-bit and a
9// 128-bit reference). Known answer (FAIL LOUD): "<ok> <total> ", ok == total.
10
11import "nx_mul_wide.nx"
12
13const T_MASK32: i64 = 4294967295
14const T_TWO32: i64 = 4294967296
15const T_P1: i64 = 2147483647
16const T_P2: i64 = 1000000007
17const T_P3: i64 = 998244353
18const T_LCG_A: i64 = 6364136223846793005
19const T_LCG_C: i64 = 1442695040888963407
20
21func _emit_num(v: i64) -> i64 {
22 let b: *u8 = sys_mmap(28); var n: i64 = v; if n < 0 { n = 0 - n }
23 let t2: *u8 = sys_mmap(28); var t: i64 = 0
24 if n == 0 { t2[0] = 48; t = 1 }
25 while n > 0 { t2[t] = 48 + (n % 10); n = n / 10; t = t + 1 }
26 var i: i64 = 0; while i < t { b[i] = t2[t - 1 - i]; i = i + 1 }
27 b[t] = 32; sys_write(1, b, t + 1); return 0
28}
29func _nl() -> i64 { let z: *u8 = sys_mmap(2); z[0] = 10; sys_write(1, z, 1); return 0 }
30
31// unsigned compare x < y (1/0)
32func u_lt(x: i64, y: i64) -> i64 {
33 if x >= 0 { if y >= 0 { if x < y { return 1 } return 0 } return 1 }
34 if y >= 0 { return 0 }
35 if x < y { return 1 } return 0
36}
37// logical a >> k for k in 1..63 (mask off the arithmetic sign-extension)
38func u_shr(a: i64, k: i64) -> i64 {
39 let keep: i64 = 64 - k
40 let mask: i64 = (1 << keep) - 1
41 return (a >> k) & mask
42}
43// unsigned x mod p (p < 2^31)
44func t_umod(x: i64, p: i64) -> i64 {
45 let xl: i64 = x & T_MASK32
46 let xh: i64 = (x >> 32) & T_MASK32
47 return ((xh % p) * (T_TWO32 % p) + (xl % p)) % p
48}
49func t_check(a: i64, b: i64, hi: i64, lo: i64, p: i64) -> i64 {
50 let lhs: i64 = (t_umod(a, p) * t_umod(b, p)) % p
51 let t32: i64 = T_TWO32 % p
52 let t64: i64 = (t32 * t32) % p
53 let rhs: i64 = ((t_umod(hi, p) * t64) % p + t_umod(lo, p)) % p
54 if lhs == rhs { return 1 }
55 return 0
56}
57// LEG B: independent 128-bit reference via bit-by-bit shift-add (two-word).
58func ref_mul128(a: i64, b: i64, hi_out: *i64) -> i64 {
59 var hi: i64 = 0
60 var lo: i64 = 0
61 var i: i64 = 0
62 while i < 64 {
63 let bit: i64 = (b >> i) & 1
64 if bit == 1 {
65 var add_lo: i64 = a
66 var add_hi: i64 = 0
67 if i > 0 { add_lo = a << i; add_hi = u_shr(a, 64 - i) }
68 let new_lo: i64 = lo + add_lo
69 var carry: i64 = 0
70 if u_lt(new_lo, lo) == 1 { carry = 1 }
71 lo = new_lo
72 hi = hi + add_hi + carry
73 }
74 i = i + 1
75 }
76 hi_out[0] = hi
77 return lo
78}
79
80func _trial(g: *NxGsim, lonet: i64, hinet: i64, a: i64, b: i64) -> i64 { // 1 if all legs agree
81 g.vals[0] = a; g.vals[1] = b
82 if nx_gsim_run(g) != NX_GSIM_OK { return 0 }
83 let loA: i64 = g.vals[lonet]
84 let hiA: i64 = g.vals[hinet]
85 let rho: *i64 = sys_mmap(8) as *i64; rho[0] = 0
86 let loB: i64 = ref_mul128(a, b, rho)
87 let hiB: i64 = rho[0]
88 if loA != a * b { return 0 } // sanity: low == i64 product
89 if loA != loB { return 0 } // LEG A vs LEG B (lo)
90 if hiA != hiB { return 0 } // LEG A vs LEG B (hi) -- independent 128-bit
91 if t_check(a, b, hiA, loA, T_P1) == 0 { return 0 } // LEG C
92 if t_check(a, b, hiA, loA, T_P2) == 0 { return 0 }
93 if t_check(a, b, hiA, loA, T_P3) == 0 { return 0 }
94 return 1
95}
96
97func main() -> i64 {
98 let vals: *i64 = sys_mmap(128 * 8) as *i64
99 let cells: *NxGsimCell = sys_mmap(128 * 48) as *NxGsimCell
100 let g: *NxGsim = sys_mmap(64) as *NxGsim
101 g.vals = vals; g.n_nets = 2; g.cells = cells; g.n_cells = 0
102 let hio: *i64 = sys_mmap(8) as *i64; hio[0] = 0
103 let lonet: i64 = nx_mul_wide_synth(g, 0, 1, hio)
104 let hinet: i64 = hio[0]
105
106 var total: i64 = 0
107 var ok: i64 = 0
108
109 let ea: *i64 = sys_mmap(8 * 6) as *i64
110 let eb: *i64 = sys_mmap(8 * 6) as *i64
111 ea[0]=0; eb[0]=0
112 ea[1]=1; eb[1]=1
113 ea[2]=0-1; eb[2]=0-1
114 ea[3]=0-1; eb[3]=1
115 ea[4]=4294967296; eb[4]=4294967296
116 ea[5]=123456789; eb[5]=987654321999
117 var e: i64 = 0
118 while e < 6 { total = total + 1; if _trial(g, lonet, hinet, ea[e], eb[e]) == 1 { ok = ok + 1 } e = e + 1 }
119
120 var seed: i64 = 88172645463325252
121 var nr: i64 = 0
122 while nr < 50000 {
123 seed = seed * T_LCG_A + T_LCG_C; let a: i64 = seed
124 seed = seed * T_LCG_A + T_LCG_C; let b: i64 = seed
125 total = total + 1
126 if _trial(g, lonet, hinet, a, b) == 1 { ok = ok + 1 }
127 nr = nr + 1
128 }
129
130 _emit_num(ok); _emit_num(total); _nl()
131 if ok != total { sys_exit(1); return 1 } // all 3 independent legs agree on every case
132 if total != 50006 { sys_exit(2); return 2 }
133 sys_exit(0); return 0
134}