File size: 7,568 Bytes
409010d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
"""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()