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}