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