| """Vanilla 3-layer feedforward ANN for modular multiplication. |
| |
| Input: digit-encoded (a_red, b_red, p) — each zero-padded to MAX_DIGITS |
| Hidden: one fully-connected layer with ReLU |
| Output: MAX_OUT_DIGITS * 10 logits — one 10-way class per output digit position |
| |
| At inference, a and b are reduced mod p inside predict_digits (allowed by rules) |
| before being encoded and fed to the network. |
| """ |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import torch |
| import torch.nn as nn |
|
|
| from modchallenge.interface.base_model import ModularMultiplicationModel |
|
|
| MAX_DIGITS = 10 |
| MAX_OUT_DIGITS = 10 |
| INPUT_SIZE = 3 * MAX_DIGITS |
| HIDDEN_SIZE = 1024 |
| OUTPUT_SIZE = MAX_OUT_DIGITS * 10 |
|
|
|
|
| |
| |
| |
|
|
| class VanillaMLP(nn.Module): |
| def __init__( |
| self, |
| input_size: int = INPUT_SIZE, |
| hidden_size: int = HIDDEN_SIZE, |
| output_size: int = OUTPUT_SIZE, |
| ): |
| super().__init__() |
| self.net = nn.Sequential( |
| nn.Linear(input_size, hidden_size), |
| nn.ReLU(), |
| nn.Linear(hidden_size, hidden_size), |
| nn.ReLU(), |
| nn.Linear(hidden_size, hidden_size), |
| nn.ReLU(), |
| nn.Linear(hidden_size, output_size), |
| ) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| """x: (B, INPUT_SIZE) floats in [0, 1]. returns (B, MAX_OUT_DIGITS, 10) logits.""" |
| out = self.net(x) |
| return out.view(-1, MAX_OUT_DIGITS, 10) |
|
|
|
|
| |
| |
| |
|
|
| def encode_int(n: int, length: int = MAX_DIGITS) -> list[int]: |
| """Integer -> zero-padded decimal digit list of fixed length, MSB first.""" |
| s = str(int(n)).zfill(length) |
| if len(s) > length: |
| s = s[-length:] |
| return [int(c) for c in s] |
|
|
|
|
| def digits_to_tensor(n: int, length: int = MAX_DIGITS) -> list[float]: |
| """Encode integer n as normalized floats in [0, 1].""" |
| return [d / 9.0 for d in encode_int(n, length)] |
|
|
|
|
| |
| |
| |
|
|
| class VanillaModel(ModularMultiplicationModel): |
| def __init__(self): |
| self.model: VanillaMLP | None = None |
| self.device: torch.device | None = None |
|
|
| def load(self, model_dir: str) -> None: |
| if 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, |
| ) |
| cfg = ckpt.get("config", {}) |
| self.model = VanillaMLP( |
| input_size=cfg.get("input_size", INPUT_SIZE), |
| hidden_size=cfg.get("hidden_size", HIDDEN_SIZE), |
| output_size=cfg.get("output_size", OUTPUT_SIZE), |
| ) |
| self.model.load_state_dict(ckpt["state_dict"]) |
| self.model.to(self.device) |
| self.model.eval() |
|
|
| def preprocess_a(self, a: str): |
| return a |
|
|
| def preprocess_b(self, b: str): |
| return b |
|
|
| def preprocess_p(self, p: str): |
| return p |
|
|
| @torch.no_grad() |
| def predict_digits(self, a_enc, b_enc, p_enc) -> list[int]: |
| assert self.model is not None |
|
|
| p = int(p_enc) |
| a_red = int(a_enc) % p |
| b_red = int(b_enc) % p |
|
|
| x = digits_to_tensor(a_red) + digits_to_tensor(b_red) + digits_to_tensor(p) |
| inp = torch.tensor([x], dtype=torch.float32, device=self.device) |
|
|
| logits = self.model(inp) |
| preds = logits[0].argmax(-1).tolist() |
|
|
| |
| result = preds |
| while len(result) > 1 and result[0] == 0: |
| result = result[1:] |
| return result |
|
|