"""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 # decimal digits per input slot (covers p up to ~4 billion = tier 4) MAX_OUT_DIGITS = 10 # decimal digits in the answer (answer < p <= tier-4 max) INPUT_SIZE = 3 * MAX_DIGITS # 30 HIDDEN_SIZE = 1024 OUTPUT_SIZE = MAX_OUT_DIGITS * 10 # 100 # --------------------------------------------------------------------------- # Architecture # --------------------------------------------------------------------------- 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) # (B, OUTPUT_SIZE) return out.view(-1, MAX_OUT_DIGITS, 10) # (B, D, 10) # --------------------------------------------------------------------------- # Encoding helpers (shared by train.py and predict_digits) # --------------------------------------------------------------------------- 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:] # truncate if somehow too long 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)] # --------------------------------------------------------------------------- # Submission entry point # --------------------------------------------------------------------------- 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) # (1, 30) logits = self.model(inp) # (1, MAX_OUT_DIGITS, 10) preds = logits[0].argmax(-1).tolist() # [d0, d1, ..., d9] # Strip leading zeros, return at least [0] result = preds while len(result) > 1 and result[0] == 0: result = result[1:] return result