File size: 4,433 Bytes
df42b8c aa11a0b df42b8c aa11a0b df42b8c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | """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
|