from __future__ import annotations import copy import random import sys import torch from model import (make_reduce_cell, make_add_cell, BitStreamMachine, _bits_of, PAD_HEAD) def probable_prime(rng: random.Random, bits: int) -> int: def is_pp(n): if n < 2: return False for sp in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31): if n % sp == 0: return n == sp d, r = n - 1, 0 while d % 2 == 0: d //= 2 r += 1 for _ in range(20): a = rng.randrange(2, n - 1) x = pow(a, d, n) if x in (1, n - 1): continue for _ in range(r - 1): x = x * x % n if x == n - 1: break else: return False return True while True: p = rng.getrandbits(bits - 1) | (1 << (bits - 1)) | 1 if is_pp(p): return p @torch.no_grad() def accuracy(cells, rng: random.Random, p_bits: int, op_bits: int, n_problems: int = 50) -> float: mach = BitStreamMachine(cells[0], cells[1], torch.device("cpu")) probs = [] for _ in range(n_problems): p = probable_prime(rng, p_bits) a = rng.getrandbits(rng.randint(1, op_bits)) b = rng.getrandbits(rng.randint(1, op_bits)) probs.append((a, b, p)) n_p = max(p.bit_length() for _, _, p in probs) + PAD_HEAD L = max(2, max(max(a.bit_length(), b.bit_length()) for a, b, _ in probs)) L += L % 2 def pack(vals, w): m = torch.zeros(len(vals), w) for r, v in enumerate(vals): bits = _bits_of(v) m[r, w - len(bits):] = torch.tensor(bits, dtype=torch.float32) return m z = mach.run(pack([a for a, _, _ in probs], L), pack([b for _, b, _ in probs], L), pack([p for _, _, p in probs], n_p), pack([3 * p for _, _, p in probs], n_p)) good = 0 for r, (a, b, p) in enumerate(probs): got = int("".join(str(int(v)) for v in z[r].tolist()), 2) good += (got == (a * b) % p) return good / n_problems def main(): ckpt_path = sys.argv[1] if len(sys.argv) > 1 else "weights.pt" ck = torch.load(ckpt_path, map_location="cpu", weights_only=True) rcell = make_reduce_cell() rcell.load_state_dict(ck["reduce_state_dict"]) rcell.eval() acell = make_add_cell() acell.load_state_dict(ck["add_state_dict"]) acell.eval() cells = (rcell, acell) rng = random.Random(42) print("trained weights:") for pb, ob in ((14, 64), (28, 96)): print(f" p ~ {pb} bits, ops {ob} bits: " f"accuracy {accuracy(cells, rng, pb, ob):.2f}") for scale in (0.02, 0.1): pert = tuple(copy.deepcopy(c) for c in cells) torch.manual_seed(0) with torch.no_grad(): for c in pert: for prm in c.parameters(): prm.add_(torch.randn_like(prm) * scale * (prm.abs().mean() + 1e-8)) print(f"weights + {scale:.0%} relative noise:") for pb, ob in ((14, 64), (28, 96)): print(f" p ~ {pb} bits, ops {ob} bits: " f"accuracy {accuracy(pert, rng, pb, ob):.2f}") torch.manual_seed(1) fresh = (make_reduce_cell().eval(), make_add_cell().eval()) print("reinitialized (untrained) weights:") for pb, ob in ((14, 64), (28, 96)): print(f" p ~ {pb} bits, ops {ob} bits: " f"accuracy {accuracy(fresh, rng, pb, ob):.2f}") if __name__ == "__main__": main()