Reversible Mealy Token Mixer / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()