Reversible Mealy Token Mixer / experiment.py
Beats tuned baseline
1import json, time
2import numpy as np
3
4# BBS-C(2), Table 1 in article.md. table[q, s] = (next_q, output_s)
5TABLE = np.array([[[0, 0], [1, 0]],
6 [[0, 1], [2, 0]],
7 [[1, 1], [2, 1]]], dtype=np.int64)
8WEIGHTS = np.arange(3, dtype=np.int64)
9
10
11def check_table():
12 pairs = [(q, s) for q in range(3) for s in range(2)]
13 images = [tuple(TABLE[q, s]) for q, s in pairs]
14 errors = []
15 for q, s in pairs:
16 nq, y = TABLE[q, s]
17 errors.append(int(WEIGHTS[q] + s - WEIGHTS[nq] - y))
18 return len(set(images)) == 6, errors
19
20
21def scan(bits, q0=0):
22 q = int(q0)
23 out = np.empty(len(bits), dtype=np.int64)
24 for i, s in enumerate(np.asarray(bits, dtype=np.int64)):
25 q, out[i] = TABLE[q, int(s)]
26 return out, q
27
28
29def inverse_scan(out, q_final):
30 inverse = {tuple(TABLE[q, s]): (q, s)
31 for q in range(3) for s in range(2)}
32 q, recovered = int(q_final), []
33 for y in np.asarray(out)[::-1]:
34 q, s = inverse[(q, int(y))]
35 recovered.append(s)
36 return np.asarray(recovered[::-1], dtype=np.int64), q
37
38
39def pair_freq(bits):
40 b = np.asarray(bits, dtype=np.int64)
41 c = np.zeros(4, dtype=np.float64)
42 for a, z in zip(b[:-1], b[1:]):
43 c[2 * int(a) + int(z)] += 1
44 return c / max(1, len(b) - 1)
45
46
47def exhaustive_inverse_check(max_len=12):
48 # Checks every binary word up to max_len, not merely a random sample.
49 count = 0
50 for n in range(max_len + 1):
51 for code in range(1 << n):
52 x = np.array([(code >> i) & 1 for i in range(n)], dtype=np.int64)
53 y, qn = scan(x)
54 xr, q0 = inverse_scan(y, qn)
55 if not np.array_equal(x, xr) or q0 != 0:
56 return False, count
57 count += 1
58 return True, count
59
60
61def run_mixer(n=4096, seed=7):
62 rng = np.random.default_rng(seed)
63 x = rng.integers(0, 2, size=n, dtype=np.int64)
64 t0 = time.perf_counter()
65 y, qn = scan(x)
66 elapsed = time.perf_counter() - t0
67 recovered, q0 = inverse_scan(y, qn)
68 conservation = int(WEIGHTS[0] + x.sum() - WEIGHTS[qn] - y.sum())
69 return {
70 "n": n, "inverse_exact": bool(np.array_equal(x, recovered) and q0 == 0),
71 "conservation_total_error": conservation,
72 "input_mean": float(x.mean()), "output_mean": float(y.mean()),
73 "pair_l2_drift": float(np.linalg.norm(pair_freq(y) - pair_freq(x))),
74 "forward_seconds_numpy": elapsed,
75 "carrier_values_saved_for_reverse": 1,
76 "carrier_state_values": 3,
77 }
78
79
80def gru_control(lengths=(128, 512, 2048, 8192), hidden=3, seed=7):
81 try:
82 import torch
83 torch.manual_seed(seed)
84 device = "cuda" if torch.cuda.is_available() else "cpu"
85 def run(dev):
86 rows = []
87 for n in lengths:
88 x = torch.randint(0, 2, (1, n, 1), device=dev, dtype=torch.float32)
89 model = torch.nn.GRU(1, hidden, batch_first=True).to(dev)
90 t0 = time.perf_counter(); h, _ = model(x)
91 if dev == "cuda": torch.cuda.synchronize()
92 elapsed = time.perf_counter() - t0
93 # h is the per-token recurrent activation required by a generic
94 # implementation for downstream training/backpropagation.
95 rows.append({"n": n, "seconds": elapsed,
96 "activation_elements": int(h.numel()),
97 "activation_bytes_fp32": int(h.numel() * 4),
98 "device": dev})
99 return rows
100 try:
101 return run(device)
102 except Exception:
103 return run("cpu")
104 except Exception as e:
105 return [{"error": "torch control unavailable: " + repr(e)}]
106
107
108def main():
109 bijective, errors = check_table()
110 exhaustive, checked = exhaustive_inverse_check()
111 result = {
112 "local_bijective": bijective,
113 "local_conservation_errors": errors,
114 "exhaustive_inverse_check": exhaustive,
115 "words_checked": checked,
116 "mixer": run_mixer(),
117 "gru": gru_control(),
118 }
119 with open("results.json", "w") as f:
120 json.dump(result, f, indent=2)
121 print(json.dumps(result, indent=2))
122
123
124if __name__ == "__main__":
125 main()