"""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 # noqa: E402 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) # Keep the random limb embeddings fixed. The trained MLPs must learn the # finite semantics from these arbitrary distributed representations. adder.limb_emb.weight.requires_grad_(False) reducer.limb_emb.weight.requires_grad_(False) # A three-symbol target alphabet. The encoders must learn to map all 1024 # limb pairs into it; the learned operators must implement its algebra. 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, ) # These parameters are part of the differentiable research modules, though # finite-table inference uses the learned codebook identities directly. 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()