Download perturb_test.py from HCho/bitstream-modmul-model: direct link, hf CLI and curl.
- Browser
- Download file 3.66 kB
-
https://huggingface.co/HCho/bitstream-modmul-model/resolve/main/perturb_test.py
- Command line
-
hf download hf://HCho/bitstream-modmul-model/perturb_test.py
-
curl -L -o perturb_test.py https://huggingface.co/HCho/bitstream-modmul-model/resolve/main/perturb_test.py
3.66 kB
| 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 | |
| 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() | |