--- license: apache-2.0 library_name: pytorch tags: - modular-arithmetic - algorithmic-reasoning - rnn - number-theory - neural-algorithm --- # Horner-RNN — learned modular multiplication up to 2⁵¹² A compliant **bit-sequential RNN** that computes `(a · b) mod p` for primes `p` up to **2⁵¹²**, by *learning the Horner step of double-and-add* rather than memorising multiplication tables. Entry for the [Modular Arithmetic Challenge](https://github.com/SAIRcompetition/modular-arithmetic-challenge). - **Saturates tiers 1–8** (all primes `< 2⁵¹²`): tiers 1–3 = 100%, tier 4 = 99%, tier 5 = 98%, tier 6 = 97%, tier 7 = 98%, **tier 8 = 92%** (512-bit) - **overall_accuracy 0.788**, `highest_tier_above_90 = 8` - The 128/256/512-bit (tier 6/7/8) cells are **carry-aware TCNs** (weight-shared dilated convolutions over the bit-positions, ~4–6M params each) — a far better inductive bias for long carry chains than the MLP, and the key to the per-step precision a 128/256/512-step chain demands. The per-step error floor rises with width, so the 512-bit cell additionally uses **gradient accumulation** (a large effective batch lowers the per-step noise floor) to reach tier 8 = 0.92 - Verifiably **generalises to primes never seen in training** (held-out-prime validation accuracy tracks training accuracy — no memorisation gap) ## The idea Write `a` in bits, MSB-first; then `a·b mod p` is the iterate of one small map: ``` t_0 = 0 t_{k+1} = (2·t_k + a_bit_k · b) mod p # one learned step (Horner) answer = t_N (N = bit width of p) ``` The model is an RNN whose **transition function — an MLP for the 16/32-bit cells, a carry-aware TCN for the 64/128/256/512-bit cells — is trained on exactly that single-step map** over binary-encoded inputs. The hidden state is a quantized bit vector (a hard binary bottleneck), so the recurrence composes cleanly: if the cell is exact per step, the chain is exact end-to-end. At inference the scan feeds the bits of `a mod p` one per step, conditioned on `(b mod p, p)`, and the final hidden-state bits are emitted MSB-first as the base-2 answer (`output_base: 2`). The single-step function is **piecewise linear** (`2t + bit·b`, then subtract `0`, `p`, or `2p`), which is why it generalises across primes where the full bilinear map `(a,b) → a·b mod p` does not. ## Files / cells The model ships **six cells** and routes each problem to the narrowest one whose state holds the prime: | File | Cell | Primes | Tiers | Arch | Params | Public benchmark | |---|---|---|---|---|---|---| | `weights16.pt` | 16-bit | `< 2¹⁶` | 1–3 | MLP, 4096 / 4 | ~50M | tiers 1–3 = 1.00 | | `weights32.pt` | 32-bit | `< 2³²` | 4 | MLP, 6144 / 4 | ~114M | tier 4 = 0.99 | | `weights64.pt` | 64-bit | `< 2⁶⁴` | 5 | **carry-aware TCN**, 256ch / 8 blocks, dilations 1–32 | ~3.2M | **tier 5 = 0.99** | | `weights128.pt` | 128-bit | `< 2¹²⁸` | 6 | **carry-aware TCN**, 256ch / 10 blocks, dilations 1–64 | ~3.9M | **tier 6 = 0.97** | | `weights256.pt` | 256-bit | `< 2²⁵⁶` | 7 | **carry-aware TCN**, 256ch / 12 blocks, dilations 1–128 | ~4.7M | **tier 7 = 0.98** | | `weights512.pt` | 512-bit | `< 2⁵¹²` | 8 | **carry-aware TCN**, 256ch / 14 blocks, dilations 1–256 | ~5.5M | **tier 8 = 0.92** | The 64/128/256/512-bit cells switch architecture: instead of a full-width MLP each is a **non-causal dilated 1-D convolutional network over the bit-positions** (64, 128, 256, 512 respectively). Carry propagation is *position-invariant* — the same carry/borrow rule applies at every bit — so a weight-shared convolution learns **one** rule applied everywhere (non-causal, so the addition carry flows LSB→MSB and the mod-`p` compare/borrow flows MSB→LSB), rather than an MLP learning a separate position-function per bit. This inductive bias drives the per-step error roughly **15× lower** than the same-task MLP — the difference between a 128/256-step chain landing at ~0.26 and at **0.97 / 0.98** — in cells **~60× smaller** than the wide MLPs (16–22 MB each vs ~950 MB). The receptive field of each TCN spans its full width in both carry directions, so a carry can propagate across the entire word. The per-step error floor *rises* with bit-width, though: the 512-bit cell needed **gradient accumulation** (a large effective batch to lower the per-step noise floor) to push its 512-step chain over the line to **tier 8 = 0.92**. The 64-bit cell is also a carry-aware TCN (it began as a 944 MB MLP that scored ~0.98 on tier 5, but had a **blind spot on primes very close to `2⁶⁴`** — the hardest top-of-range reduction — where it got 0/10; the position-invariant conv generalises to that regime, scoring 10/10, at **1/70th the size**, ~13 MB). For `p ≥ 2⁵¹²` (wider than the widest trained cell) the model emits the honest `[0]` fallback without invoking the network. Also in the repo: `model.py` (the `HornerRNN` entry class + `HornerCell`), `manifest.json` (challenge manifest), `train.py` (the 16-bit trainer). ## Usage This is a challenge submission; the base class lives in the challenge package, so install it first: ```bash pip install "git+https://github.com/SAIRcompetition/modular-arithmetic-challenge" ``` Direct inference: ```python import torch from model import HornerRNN # model.py from this repo m = HornerRNN() m.load(".") # auto-loads weights{16,32,64,128,256,512}.pt from this dir # returns base-2 digits, MSB-first; the harness decodes them to the integer digits = m.predict_digits_batch([(123456789, 987654321, 4294967291)])[0] answer = int("".join(map(str, digits)), 2) print(answer) # == (123456789 * 987654321) % 4294967291 ``` Or score it with the official harness: ```bash modchallenge evaluate . --total 1100 ``` ## Compliance (the rules permit *learned* algorithms, not hand-coded ones) The **scan** (tokenise `a mod p` into bits, iterate, read out the final state) is architecture — it computes nothing by itself. The **arithmetic** (doubling, conditional add, compare-against-`p`, carries) all lives in the trained cell weights; nothing in the code adds, multiplies, or compares against `p`. **Principle 2, measured** — perturbing the cell weights with Gaussian noise scaled to each tensor's std collapses accuracy toward the floor, and a fully re-initialised (untrained) cell is *at* the floor. The capability therefore resides in the trained parameters: | noise σ (×param std) | 0 | 0.05 | 0.1 | 0.25 | 0.5 | untrained | |---|---|---|---|---|---|---| | tier 3 (16-bit cell) | 1.00 | 1.00 | 0.98 | 0.74 | 0.06 | 0.00 | | tier 4 (32-bit cell) | 0.99 | 0.99 | 0.86 | 0.04 | 0.02 | 0.00 | | tier 5 (64-bit TCN) | 0.99 | 0.99 | 0.98 | 0.04 | 0.03 | 0.00 | | tier 6 (128-bit TCN) | 0.97 | 0.96 | 0.98 | 0.19 | 0.02 | 0.00 | | tier 7 (256-bit TCN) | 0.98 | 0.97 | 0.99 | 0.06 | 0.02 | 0.00 | | tier 8 (512-bit TCN) | 0.92 | 0.91 | 0.77 | 0.04 | 0.03 | 0.00 | Generalisation against memorisation: 10% of primes at each bit-width were held out of training entirely; chain accuracy on them matches the training primes. ## Training Single-step examples `(t, bit, b, p) → (2t + bit·b) mod p` over each tier's prime range; BCE per state bit, AdamW + cosine decay + EMA, checkpointed by full-chain accuracy on held-out primes. The TCN cells (64/128/256/512-bit) are trained on single steps drawn from the **true Horner trajectory** — `t` is an actual chain intermediate `(a_{≥i}·b) mod p`, not a uniform sample — matching the training distribution to the states the chain visits at inference (ordinary supervised BCE on the same single-step target, no backprop through the recurrence). The 128-bit (tier-6) cell is the first **carry-aware TCN**, trained over a high-diversity pool of thousands of distinct 124–128 bit primes; its weight-shared dilated-convolution bias reaches a per-step error ~15× lower than the same-task MLP, giving **tier 6 = 0.97** in a single short run. The 256-bit (tier-7) cell is the same carry-aware TCN scaled to 256 bit-positions (dilations cycling 1–128), trained identically on true-trajectory single steps over distinct 252–256 bit primes; its per-step error is low enough that the 256-step chain holds at **tier 7 = 0.98**. The 512-bit (tier-8) cell is again the same TCN (dilations 1–256) trained on distinct 510–512 bit primes; because the per-step error floor rises with width, it additionally uses **gradient accumulation** (a large effective batch lowers the gradient-noise floor on the per-step error without extra memory), which drives the 512-step chain to **tier 8 = 0.92**. Training code and the full write-up live in the solutions repo (link in the model card metadata / challenge leaderboard). ## License Apache-2.0, matching the challenge.