File size: 5,965 Bytes
bc52904
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4fce513
 
 
bc52904
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""T2-focused learned residue classifier for modular multiplication.

The model deliberately targets Tiers 1 and 2, where p < 256. It uses the same
allowed input normalization as the reference neural baselines: each operand is
reduced separately modulo p before entering the network. The network then has
to choose the output residue from learned parameters.

There is no inference-time code path that computes ``(a * b) % p``. The only
post-processing is masking classes outside ``[0, p)`` so that the emitted
single base-p digit is well-formed under the challenge decoder.
"""

from __future__ import annotations

from pathlib import Path

import torch
import torch.nn as nn

from modchallenge.interface.base_model import ModularMultiplicationModel

MAX_P = 256
MAX_CLASSES = 256
PAIR_VOCAB = MAX_P * MAX_CLASSES


class ResidueProductNet(nn.Module):
    def __init__(
        self,
        d_model: int = 128,
        hidden: int = 512,
        depth: int = 3,
        bilinear_dim: int = 64,
    ):
        super().__init__()
        self.in_emb = nn.Embedding(PAIR_VOCAB, d_model)
        self.p_emb = nn.Embedding(MAX_P, d_model)
        self.out_emb = nn.Embedding(PAIR_VOCAB, d_model)
        self.out_bias = nn.Embedding(PAIR_VOCAB, 1)
        self.left_factor = nn.Embedding(PAIR_VOCAB, bilinear_dim)
        self.right_factor = nn.Embedding(PAIR_VOCAB, bilinear_dim)
        self.candidate_factor = nn.Embedding(PAIR_VOCAB, bilinear_dim)
        self.factor_ln = nn.LayerNorm(bilinear_dim)
        self.factor_scale = bilinear_dim ** -0.5
        nn.init.zeros_(self.candidate_factor.weight)

        layers: list[nn.Module] = []
        in_dim = 4 * d_model
        for _ in range(depth):
            layers.extend(
                [
                    nn.Linear(in_dim, hidden),
                    nn.GELU(),
                    nn.LayerNorm(hidden),
                ]
            )
            in_dim = hidden
        layers.append(nn.Linear(hidden, d_model))
        layers.append(nn.LayerNorm(d_model))
        self.net = nn.Sequential(*layers)
        self.config = {
            "d_model": d_model,
            "hidden": hidden,
            "depth": depth,
            "bilinear_dim": bilinear_dim,
        }

        self.register_buffer(
            "classes", torch.arange(MAX_CLASSES, dtype=torch.long), persistent=False
        )

    def forward(self, a_red: torch.Tensor, b_red: torch.Tensor, p: torch.Tensor) -> torch.Tensor:
        a_idx = p * MAX_CLASSES + a_red
        b_idx = p * MAX_CLASSES + b_red

        ea = self.in_emb(a_idx)
        eb = self.in_emb(b_idx)
        ep = self.p_emb(p)
        h = self.net(torch.cat([ea, eb, ea * eb, ep], dim=-1))

        candidate_idx = p.unsqueeze(1) * MAX_CLASSES + self.classes.unsqueeze(0)
        candidate_emb = self.out_emb(candidate_idx)
        logits = torch.einsum("bd,bkd->bk", h, candidate_emb)
        logits = logits + self.out_bias(candidate_idx).squeeze(-1)

        # Learned low-rank residue-product factorization. This is another
        # trained head, not arithmetic post-processing: with random factors it
        # contributes no useful modular multiplication signal.
        factor_h = self.factor_ln(self.left_factor(a_idx) * self.right_factor(b_idx))
        factor_candidates = self.candidate_factor(candidate_idx)
        logits = logits + self.factor_scale * torch.einsum(
            "bd,bkd->bk", factor_h, factor_candidates
        )

        invalid = self.classes.unsqueeze(0) >= p.unsqueeze(1)
        return logits.masked_fill(invalid, -1.0e9)


class T2ResidueClassifier(ModularMultiplicationModel):
    def __init__(self):
        self.model: ResidueProductNet | None = None
        self.device: torch.device | None = None

    def load(self, model_dir: str) -> None:
        if torch.backends.mps.is_available():
            self.device = torch.device("mps")
        elif torch.cuda.is_available():
            self.device = torch.device("cuda")
        else:
            self.device = torch.device("cpu")

        ckpt = torch.load(
            Path(model_dir) / "weights.pt",
            map_location=self.device,
            weights_only=True,
        )
        self.model = ResidueProductNet(**ckpt.get("config", {}))
        load_result = self.model.load_state_dict(ckpt["state_dict"], strict=False)
        extra_keys = load_result[1]
        if extra_keys:
            raise RuntimeError(f"extra checkpoint keys: {extra_keys}")
        self.model.to(self.device)
        self.model.eval()

    def preprocess_a(self, a):
        return a

    def preprocess_b(self, b):
        return b

    def preprocess_p(self, p):
        return p

    @torch.no_grad()
    def predict_digits(self, a_enc, b_enc, p_enc):
        return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0]

    @torch.no_grad()
    def predict_digits_batch(self, inputs):
        assert self.model is not None
        assert self.device is not None

        out: list[list[int] | None] = [None] * len(inputs)
        a_rows: list[int] = []
        b_rows: list[int] = []
        p_rows: list[int] = []
        idx: list[int] = []

        for i, (a_enc, b_enc, p_enc) in enumerate(inputs):
            p = int(p_enc)
            if not (2 <= p < MAX_P):
                out[i] = [0]
                continue
            a_rows.append(int(a_enc) % p)
            b_rows.append(int(b_enc) % p)
            p_rows.append(p)
            idx.append(i)

        if idx:
            a_t = torch.tensor(a_rows, dtype=torch.long, device=self.device)
            b_t = torch.tensor(b_rows, dtype=torch.long, device=self.device)
            p_t = torch.tensor(p_rows, dtype=torch.long, device=self.device)
            preds = self.model(a_t, b_t, p_t).argmax(dim=-1).tolist()
            for j, i in enumerate(idx):
                out[i] = [int(preds[j])]

        return [row if row is not None else [0] for row in out]

    def max_batch_size(self) -> int:
        return 4096