CoolFace
Modelpublic

slkML/ModularArithmeticChallenge

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
model.py130 linesDownload Raw Back to root
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