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