| """Reproduce the submitted finite-cell training stage. |
| |
| This trains the submitted architectures from random initialization on the |
| complete primitive domains, writes a fresh artifact, and refuses to finish |
| until every extracted table cell matches its specification. |
| |
| The arithmetic below is used only to generate training labels and audit a new |
| checkpoint. This file is included for provenance but is never imported by the |
| inference entrypoint. |
| |
| python scripts/retrain_finite_cells.py --output reproduced_checkpoint |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import sys |
| from pathlib import Path |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
|
|
| HERE = Path(__file__).resolve().parent |
| SUBMISSION = HERE if (HERE / "model.py").is_file() else HERE.parent / "submission" |
| ROOT = SUBMISSION.parent |
| sys.path.insert(0, str(SUBMISSION)) |
|
|
| from model import BASE, ModReduce, ScanAdder, _nearest_class |
|
|
|
|
| def parameters(*modules: torch.nn.Module): |
| return [p for module in modules for p in module.parameters()] |
|
|
|
|
| def fit( |
| name: str, |
| modules: tuple[torch.nn.Module, ...], |
| loss_fn, |
| exact_fn, |
| *, |
| steps: int, |
| lr: float, |
| ) -> None: |
| optimizer = torch.optim.AdamW(parameters(*modules), lr=lr, weight_decay=1e-5) |
| for step in range(1, steps + 1): |
| optimizer.zero_grad(set_to_none=True) |
| loss = loss_fn() |
| loss.backward() |
| optimizer.step() |
| if step % 50 == 0 and exact_fn(): |
| print(f"{name:22s} exact after {step:4d} steps (loss={loss.item():.3g})") |
| return |
| raise RuntimeError(f"{name} did not become exact in {steps} steps") |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument( |
| "--output", type=Path, default=ROOT / "reproduced_checkpoint" |
| ) |
| parser.add_argument("--seed", type=int, default=20260810) |
| parser.add_argument("--steps", type=int, default=5000) |
| parser.add_argument("--lr", type=float, default=3e-3) |
| parser.add_argument("--device", choices=("cpu", "mps", "cuda")) |
| args = parser.parse_args() |
|
|
| torch.manual_seed(args.seed) |
| device = args.device or ( |
| "cuda" if torch.cuda.is_available() |
| else "mps" if torch.backends.mps.is_available() |
| else "cpu" |
| ) |
| adder = ScanAdder().to(device) |
| reducer = ModReduce().to(device) |
|
|
| |
| |
| adder.limb_emb.weight.requires_grad_(False) |
| reducer.limb_emb.weight.requires_grad_(False) |
|
|
| |
| |
| book = torch.zeros(3, 16, device=device) |
| book[:, :3] = 4 * torch.eye(3, device=device) |
|
|
| digit = torch.arange(BASE, device=device) |
| x = digit.repeat_interleave(BASE) |
| y = digit.repeat(BASE) |
| pair_a = torch.cat([adder.limb_emb(x), adder.limb_emb(y)], dim=-1) |
| pair_r = torch.cat([reducer.limb_emb(x), reducer.limb_emb(y)], dim=-1) |
| carry_label = torch.where(x + y >= BASE, 2, torch.where(x + y == BASE - 1, 1, 0)) |
| cmp_label = torch.where(x < y, 0, torch.where(x == y, 1, 2)) |
| borrow_label = torch.where(x < y, 2, torch.where(x == y, 1, 0)) |
|
|
| fit( |
| "adder encoder", |
| (adder.encoder,), |
| lambda: F.mse_loss(adder.encoder(pair_a), book[carry_label]), |
| lambda: bool((_nearest_class(adder.encoder(pair_a), book) == carry_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
| fit( |
| "comparison encoder", |
| (reducer.cmp_encoder,), |
| lambda: F.mse_loss(reducer.cmp_encoder(pair_r), book[cmp_label]), |
| lambda: bool((_nearest_class(reducer.cmp_encoder(pair_r), book) == cmp_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
| fit( |
| "borrow encoder", |
| (reducer.borrow_encoder,), |
| lambda: F.mse_loss(reducer.borrow_encoder(pair_r), book[borrow_label]), |
| lambda: bool((_nearest_class(reducer.borrow_encoder(pair_r), book) == borrow_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
|
|
| left = torch.arange(3, device=device).repeat_interleave(3) |
| right = torch.arange(3, device=device).repeat(3) |
| op_input = torch.cat([book[left], book[right]], dim=-1) |
| carry_op_label = torch.where(right == 1, left, right) |
| cmp_op_label = torch.where(left == 1, right, left) |
| fit( |
| "adder operator", |
| (adder.op,), |
| lambda: F.mse_loss(adder.op(op_input), book[carry_op_label]), |
| lambda: bool((_nearest_class(adder.op(op_input), book) == carry_op_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
| fit( |
| "comparison operator", |
| (reducer.cmp_op,), |
| lambda: F.mse_loss(reducer.cmp_op(op_input), book[cmp_op_label]), |
| lambda: bool((_nearest_class(reducer.cmp_op(op_input), book) == cmp_op_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
| fit( |
| "borrow operator", |
| (reducer.borrow_op,), |
| lambda: F.mse_loss(reducer.borrow_op(op_input), book[carry_op_label]), |
| lambda: bool((_nearest_class(reducer.borrow_op(op_input), book) == carry_op_label).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
|
|
| prefix = torch.arange(3, device=device) |
| pair_a3 = pair_a[:, None, :].expand(BASE * BASE, 3, -1) |
| prefix3 = book[None, :, :].expand(BASE * BASE, 3, -1) |
| adder_input = torch.cat([pair_a3, prefix3], dim=-1).reshape(-1, 80) |
| carry_in = (prefix == 2).long() |
| adder_target = ((x[:, None] + y[:, None] + carry_in) % BASE).reshape(-1) |
| fit( |
| "adder resolver", |
| (adder.resolver,), |
| lambda: F.cross_entropy(adder.resolver(adder_input), adder_target), |
| lambda: bool((adder.resolver(adder_input).argmax(-1) == adder_target).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
|
|
| verdict = torch.arange(3, device=device) |
| pair_r9 = pair_r[:, None, None, :].expand(BASE * BASE, 3, 3, -1) |
| borrow9 = book[None, :, None, :].expand(BASE * BASE, 3, 3, -1) |
| verdict9 = book[None, None, :, :].expand(BASE * BASE, 3, 3, -1) |
| reducer_input = torch.cat([pair_r9, borrow9, verdict9], dim=-1).reshape(-1, 96) |
| subtract_digit = ( |
| x[:, None, None] - y[:, None, None] - carry_in[None, :, None] |
| ) % BASE |
| reducer_target = torch.where( |
| verdict[None, None, :] == 0, |
| x[:, None, None], |
| subtract_digit, |
| ).expand(-1, 3, 3).reshape(-1) |
| fit( |
| "reducer resolver", |
| (reducer.resolver,), |
| lambda: F.cross_entropy(reducer.resolver(reducer_input), reducer_target), |
| lambda: bool((reducer.resolver(reducer_input).argmax(-1) == reducer_target).all()), |
| steps=args.steps, lr=args.lr, |
| ) |
|
|
| |
| |
| with torch.no_grad(): |
| adder.identity.copy_(book[1]) |
| reducer.cmp_identity.copy_(book[1]) |
| reducer.borrow_identity.copy_(book[1]) |
|
|
| args.output.mkdir(parents=True, exist_ok=True) |
| torch.save({k: v.detach().cpu() for k, v in adder.state_dict().items()}, args.output / "adder.pt") |
| torch.save({k: v.detach().cpu() for k, v in reducer.state_dict().items()}, args.output / "reducer.pt") |
| torch.save( |
| {"carry": book.cpu(), "cmp": book.cpu(), "borrow": book.cpu()}, |
| args.output / "codebooks.pt", |
| ) |
| print(f"wrote clean-room checkpoint to {args.output}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|