slkML's picture
updated model.py
aa11a0b verified
Raw
History Blame Contribute Delete
4.43 kB
"""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