from __future__ import annotations import argparse import random import sys import time from pathlib import Path import torch sys.path.insert(0, str(Path(__file__).resolve().parent)) from model import ( make_reduce_cell, make_add_cell, reduce_features, add_features, BitStreamMachine, _bits_of, PAD_HEAD, ) from data import make_reduce_batch, make_add_batch def is_probable_prime(n: int, rng: random.Random, rounds: int = 24) -> bool: if n < 2: return False for sp in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37): if n % sp == 0: return n == sp d, r = n - 1, 0 while d % 2 == 0: d //= 2 r += 1 for _ in range(rounds): 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 def random_prime(rng: random.Random, lo_bits: int, hi_bits: int) -> int: while True: l = rng.randint(lo_bits, hi_bits) if l == 1: return 2 p = rng.getrandbits(l - 1) | (1 << (l - 1)) | 1 if p == 1: continue if is_probable_prime(p, rng): return p def load_cells(path: str): ck = torch.load(path, map_location="cpu", weights_only=True) rcell = make_reduce_cell() rcell.load_state_dict(ck.get("reduce_ema_state_dict", ck["reduce_state_dict"])) rcell.eval() acell = make_add_cell() acell.load_state_dict(ck.get("add_ema_state_dict", ck["add_state_dict"])) acell.eval() return rcell, acell @torch.no_grad() def cell_stress(rcell, acell, rng, widths, per_width): print("== cell-level stress (hard gates, CPU) ==") worst = 1.0 for n in widths: bsz = max(8, 20000 // n) for kind, cell in (("reduce", rcell), ("add", acell)): total, good = 0, 0 while total < per_width: if kind == "reduce": b = make_reduce_batch(rng, n, bsz) feats = reduce_features(b["x"], b["p"], b["p3"]) else: b = make_add_batch(rng, n, bsz) feats = add_features(b["x"], b["y"], b["g"]) ok = ((cell(feats) > 0) == (b["z"] > 0.5)).all(dim=1) good += int(ok.sum()) total += ok.numel() rate = good / total worst = min(worst, rate) print(f" n={n:5d} {kind:6s}: {good}/{total} ({rate:.6f})" + (" <-- FAILURES" if rate < 1 else "")) return worst 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 @torch.no_grad() def machine_tier_check(rcell, acell, rng, per_tier, tiers=None): print("== machine-level end-to-end vs ground truth ==") mach = BitStreamMachine(rcell, acell, torch.device("cpu")) geo = {1: (1, 3, 32), 2: (4, 8, 48), 3: (9, 16, 64), 4: (17, 32, 96), 5: (33, 64, 128), 6: (65, 128, 256), 7: (129, 256, 512), 8: (257, 512, 1024), 9: (513, 1024, 2048)} if tiers: geo = {t: geo[t] for t in tiers} results = {} for t, (lo, hi, ob) in geo.items(): primes = [random_prime(rng, lo, hi) for _ in range(5)] probs = [] for i in range(per_tier): p = primes[i % 5] if i < 4: a, b = [(0, rng.getrandbits(ob)), (rng.getrandbits(ob), 0), (1, rng.getrandbits(ob)), (rng.getrandbits(ob), 1)][i] else: a = rng.getrandbits(rng.randint(1, ob)) b = rng.getrandbits(rng.randint(1, ob)) 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 t0 = time.time() 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)) dt = time.time() - t0 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) results[t] = (good, per_tier, dt) print(f" tier {t}: {good}/{per_tier} exact ({dt:.1f}s for batch)") return results @torch.no_grad() def batch_invariance(rcell, acell, rng): print("== batch-composition invariance ==") mach = BitStreamMachine(rcell, acell, torch.device("cpu")) p = random_prime(rng, 60, 64) a, b = rng.getrandbits(128), rng.getrandbits(128) def run_with(width_pad, extra_p_bits): n_p = p.bit_length() + PAD_HEAD + extra_p_bits L = 128 + width_pad L += L % 2 z = mach.run(_pack([a], L), _pack([b], L), _pack([p], n_p), _pack([3 * p], n_p)) return int("".join(str(int(v)) for v in z[0].tolist()), 2) vals = {run_with(wp, ep) for wp in (0, 8, 32) for ep in (0, 5, 40)} ok = len(vals) == 1 and vals == {(a * b) % p} print(f" distinct outputs across paddings: {len(vals)} (want 1), " f"correct={ok}") return ok @torch.no_grad() def perturbation(rcell, acell, rng): print("== weight-perturbation collapse (compliance evidence) ==") import copy for scale in (0.0, 0.02, 0.1): r2, m2 = copy.deepcopy(rcell), copy.deepcopy(acell) if scale: for c in (r2, m2): for prm in c.parameters(): prm.add_(torch.randn_like(prm) * scale * (prm.abs().mean() + 1e-8)) res = machine_tier_check(r2, m2, random.Random(7), 20, tiers=[3]) g, n, _ = res[3] print(f" noise scale {scale}: tier-3 accuracy {g}/{n}") def main(): ap = argparse.ArgumentParser() ap.add_argument("ckpt") ap.add_argument("--quick", action="store_true") ap.add_argument("--tiers", type=int, nargs="*", default=None) ap.add_argument("--per-tier", type=int, default=100) ap.add_argument("--per-width", type=int, default=200_000) ap.add_argument("--seed", type=int, default=123) ap.add_argument("--skip-stress", action="store_true") args = ap.parse_args() torch.set_num_threads(8) rng = random.Random(args.seed) rcell, acell = load_cells(args.ckpt) widths = [5, 6, 8, 11, 15, 19, 23, 27, 31, 35, 36, ] per_width = 20_000 if args.quick else args.per_width per_tier = 30 if args.quick else args.per_tier worst = 1.0 if not args.skip_stress: worst = cell_stress(rcell, acell, rng, widths, per_width) machine_tier_check(rcell, acell, rng, per_tier, tiers=args.tiers) inv = batch_invariance(rcell, acell, rng) if not args.quick: perturbation(rcell, acell, rng) print(f"\nworst cell width rate: {worst:.6f}; batch-invariant: {inv}") if __name__ == "__main__": main()