uwunion's picture
Publish uniform-transition competition artifact
0a84378 verified
Raw
History Blame Contribute Delete
7.57 kB
"""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()