slkML/ModularArithmeticChallenge
0
1"""Vanilla 3-layer feedforward ANN for modular multiplication.2 3Input: digit-encoded (a_red, b_red, p) — each zero-padded to MAX_DIGITS4Hidden: one fully-connected layer with ReLU5Output: MAX_OUT_DIGITS * 10 logits — one 10-way class per output digit position6 7At inference, a and b are reduced mod p inside predict_digits (allowed by rules)8before being encoded and fed to the network.9"""10 11from __future__ import annotations12 13from pathlib import Path14 15import torch16import torch.nn as nn17 18from modchallenge.interface.base_model import ModularMultiplicationModel19 20MAX_DIGITS = 10 # decimal digits per input slot (covers p up to ~4 billion = tier 4)21MAX_OUT_DIGITS = 10 # decimal digits in the answer (answer < p <= tier-4 max)22INPUT_SIZE = 3 * MAX_DIGITS # 3023HIDDEN_SIZE = 102424OUTPUT_SIZE = MAX_OUT_DIGITS * 10 # 10025 26 27# ---------------------------------------------------------------------------28# Architecture29# ---------------------------------------------------------------------------30 31class VanillaMLP(nn.Module):32 def __init__(33 self,34 input_size: int = INPUT_SIZE,35 hidden_size: int = HIDDEN_SIZE,36 output_size: int = OUTPUT_SIZE,37 ):38 super().__init__()39 self.net = nn.Sequential(40 nn.Linear(input_size, hidden_size),41 nn.ReLU(),42 nn.Linear(hidden_size, hidden_size),43 nn.ReLU(),44 nn.Linear(hidden_size, hidden_size),45 nn.ReLU(),46 nn.Linear(hidden_size, output_size),47 )48 49 def forward(self, x: torch.Tensor) -> torch.Tensor:50 """x: (B, INPUT_SIZE) floats in [0, 1]. returns (B, MAX_OUT_DIGITS, 10) logits."""51 out = self.net(x) # (B, OUTPUT_SIZE)52 return out.view(-1, MAX_OUT_DIGITS, 10) # (B, D, 10)53 54 55# ---------------------------------------------------------------------------56# Encoding helpers (shared by train.py and predict_digits)57# ---------------------------------------------------------------------------58 59def encode_int(n: int, length: int = MAX_DIGITS) -> list[int]:60 """Integer -> zero-padded decimal digit list of fixed length, MSB first."""61 s = str(int(n)).zfill(length)62 if len(s) > length:63 s = s[-length:] # truncate if somehow too long64 return [int(c) for c in s]65 66 67def digits_to_tensor(n: int, length: int = MAX_DIGITS) -> list[float]:68 """Encode integer n as normalized floats in [0, 1]."""69 return [d / 9.0 for d in encode_int(n, length)]70 71 72# ---------------------------------------------------------------------------73# Submission entry point74# ---------------------------------------------------------------------------75 76class VanillaModel(ModularMultiplicationModel):77 def __init__(self):78 self.model: VanillaMLP | None = None79 self.device: torch.device | None = None80 81 def load(self, model_dir: str) -> None:82 if torch.cuda.is_available():83 self.device = torch.device("cuda")84 else:85 self.device = torch.device("cpu")86 87 ckpt = torch.load(88 Path(model_dir) / "weights.pt",89 map_location=self.device,90 weights_only=True,91 )92 cfg = ckpt.get("config", {})93 self.model = VanillaMLP(94 input_size=cfg.get("input_size", INPUT_SIZE),95 hidden_size=cfg.get("hidden_size", HIDDEN_SIZE),96 output_size=cfg.get("output_size", OUTPUT_SIZE),97 )98 self.model.load_state_dict(ckpt["state_dict"])99 self.model.to(self.device)100 self.model.eval()101 102 def preprocess_a(self, a: str):103 return a104 105 def preprocess_b(self, b: str):106 return b107 108 def preprocess_p(self, p: str):109 return p110 111 @torch.no_grad()112 def predict_digits(self, a_enc, b_enc, p_enc) -> list[int]:113 assert self.model is not None114 115 p = int(p_enc)116 a_red = int(a_enc) % p117 b_red = int(b_enc) % p118 119 x = digits_to_tensor(a_red) + digits_to_tensor(b_red) + digits_to_tensor(p)120 inp = torch.tensor([x], dtype=torch.float32, device=self.device) # (1, 30)121 122 logits = self.model(inp) # (1, MAX_OUT_DIGITS, 10)123 preds = logits[0].argmax(-1).tolist() # [d0, d1, ..., d9]124 125 # Strip leading zeros, return at least [0]126 result = preds127 while len(result) > 1 and result[0] == 0:128 result = result[1:]129 return result130 