File size: 10,381 Bytes
14bef4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5052e8b
 
 
 
 
14bef4a
 
 
25ff745
 
 
 
 
 
 
 
14bef4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
25ff745
 
 
 
 
 
 
14bef4a
 
5052e8b
 
 
14bef4a
5052e8b
 
14bef4a
 
 
 
 
 
 
 
 
 
 
 
25ff745
 
 
 
 
14bef4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
"""Challenge-interface wrapper for the modmul BP-install model.

Implements ModularMultiplicationModel. Inference:
  1. per-argument: parse a, b, p to ints (each hook sees only its own arg)
  2. predict_digits: reduce x=a%p, y=b%p (legal two-operand reduction), and if
     p is within the model's trained width, greedy-decode the reverse-LSB answer
     digits from the network; otherwise emit [0] (honest out-of-regime fallback).

The answer comes entirely from the trained network on in-regime primes:
randomising the weights collapses accuracy. output_base = 10.
"""
from __future__ import annotations

import sys
from pathlib import Path

import torch

sys.path.insert(0, str(Path(__file__).resolve().parent))
# Self-contained: ByteGPT + build_model + the entmax15 recipe are VENDORED into
# this directory (coppola_pretrain_tiny.py, coppola_pretraining.py,
# train_arith_bp_supervised.py), so the submission loads with only `torch`
# available and read access limited to its own dir (the eval sandbox contract).
# entmax15 falls back to a local forward-exact impl when the `entmax` pip
# package is absent.
from train_arith_bp_supervised import TrainConfig, build_model  # noqa: E402
import encoding as enc  # noqa: E402
import composed_encoding as cenc  # noqa: E402  (composed multiply+reduce CoT)
import kvgen  # noqa: E402  (KV-cached generation; falls back to naive loop)

from modchallenge.interface.base_model import ModularMultiplicationModel  # noqa: E402


class ModMulBP(ModularMultiplicationModel):
    def __init__(self):
        self.model = None
        self.W = 1
        self.device = None
        self.regime = 10  # max p exclusive = 10**W
        self.scratchpad = False
        self.school = False

    def load(self, model_dir: str, weights: str = "weights.pt") -> None:
        torch.manual_seed(0)  # determinism is the model's responsibility (rules)
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        ckpt = torch.load(Path(model_dir) / weights,
                          map_location=self.device, weights_only=False)
        tc = TrainConfig(**ckpt["config"])
        self.model = build_model(tc, self.device)
        self.model.load_state_dict(ckpt["state_dict"])
        self.model.eval()
        self.W = ckpt["W"]
        self.scratchpad = bool(ckpt.get("scratchpad", False))
        self.school = bool(ckpt.get("school", False))
        self.composed = bool(ckpt.get("composed", False))
        if self.composed:
            self.base = ckpt["base"]
            self.scratch = bool(ckpt.get("scratch", False))
            self.cursor = bool(ckpt.get("cursor", False))
            self.subpad = bool(ckpt.get("subpad", False))
            self.stepidx = bool(ckpt.get("stepidx", False))
            self.skiptriv = bool(ckpt.get("skiptriv", False))
            self.subnum = bool(ckpt.get("subnum", False))
            self.bemit = bool(ckpt.get("bemit", False))
            self.srt = bool(ckpt.get("srt", False))
            self.regime = self.base ** self.W
        else:
            self.regime = 10 ** self.W
        self.mulonly = bool(ckpt.get("mulonly", False))
        if self.mulonly:
            # Tier-0 member: pure multiplication. p never enters the trace --
            # the operand width is the only constraint, so claim every p the
            # specialists don't (router sorts by regime, this sorts last).
            self.regime = 10 ** (2 * self.W)
            self.p_lo, self.p_hi = 2, 2 ** 4096
            return
        # Trained prime span (for router dispatch): derived from the ckpt's
        # tier list via the official tier geometry; fallback = full regime.
        self.p_lo, self.p_hi = 2, self.regime - 1
        try:
            from modchallenge.config import TIERS
            spans = [(2 ** TIERS[t].min_bits, 2 ** TIERS[t].max_bits - 1)
                     for t in ckpt.get("tiers", [])]
            if spans:
                self.p_lo = min(lo for lo, _ in spans)
                self.p_hi = min(max(hi for _, hi in spans), self.regime - 1)
        except Exception:
            pass

    # per-argument preprocessing (each sees only its own argument)
    def preprocess_a(self, a: str) -> int:
        return int(a)

    def preprocess_b(self, b: str) -> int:
        return int(b)

    def preprocess_p(self, p: str) -> int:
        return int(p)

    @torch.no_grad()
    def predict_digits(self, a_enc, b_enc, p_enc):
        return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0]

    @torch.no_grad()
    def predict_digits_batch(self, inputs):
        out = [[0]] * len(inputs)
        if self.mulonly:
            prompt_fn = lambda p, x, y, W: cenc.prompt_str(  # noqa: E731
                x, y, p, W, 10, mulonly=True)
            n_gen = cenc.gen_len(self.W, mulonly=True)
            decode_fn = lambda g, W: cenc.decode_answer(     # noqa: E731
                g, W, 10, mulonly=True)
        elif self.composed:
            prompt_fn = lambda p, x, y, W: cenc.prompt_str(  # noqa: E731
                x, y, p, W, self.base, self.subpad)
            n_gen = cenc.gen_len(self.W, self.scratch, self.cursor, self.subpad,
                                 self.stepidx, self.skiptriv, self.subnum,
                                 self.bemit, self.srt)
            decode_fn = lambda g, W: cenc.decode_answer(     # noqa: E731
                g, W, self.base, self.scratch, self.cursor, self.subpad,
                self.stepidx, self.skiptriv, self.subnum, self.bemit, self.srt)
        else:
            cot = self.scratchpad or self.school
            prompt_fn = enc.prompt_str_sp if cot else enc.prompt_str
            n_gen = (enc.school_gen_len(self.W) if self.school
                     else enc.scratchpad_len(self.W) if self.scratchpad
                     else enc.answer_len(self.W))
            decode_fn = (enc.decode_answer_school if self.school
                         else enc.decode_answer_sp if self.scratchpad
                         else enc.decode_answer)
        prompts, idx = [], []
        for i, (a, b, p) in enumerate(inputs):
            x, y = a % p, b % p
            if self.mulonly:
                if max(x, y) >= 10 ** self.W:    # operands don't fit -> honest 0
                    continue
            elif p >= self.regime:        # outside trained width -> honest 0
                continue
            prompts.append(prompt_fn(p, x, y, self.W))
            idx.append((i, x, y, p))
        if not prompts:
            return out
        # All in-regime prompts share the same length (fixed width) -> batchable.
        # latin1: composed prompts carry limb bytes >127 (limbs.py codec).
        ids = torch.tensor([list(s.encode("latin1")) for s in prompts],
                           dtype=torch.long, device=self.device)
        plen = len(prompts[0])
        # KV-cached generation: the naive loop re-forwards the whole prefix per
        # token, which blows the 5-min/1100-problem budget on long CoTs. kvgen
        # is the same computation on the same weights, token-identical
        # (validated); fall back to the naive loop on unsupported configs.
        try:
            gens = kvgen.generate_kv(self.model, ids, n_gen).tolist()
        except AssertionError:
            seq_cap = self.model.config.seq_len
            for _k in range(n_gen):
                logits, _ = self.model(ids[:, -seq_cap:])
                nxt = logits[:, -1].argmax(dim=-1, keepdim=True)
                ids = torch.cat([ids, nxt], dim=1)
            gens = ids[:, plen:].tolist()
        for row, (i, x, y, p) in zip(gens, idx):
            gen = bytes(b & 0xFF for b in row).decode("latin1")
            # No arithmetic touch-up of the model's answer: a decoded value >= p
            # would be malformed (scored incorrect anyway), so emit the honest
            # [0] fallback instead of clamping with % p.
            ans = decode_fn(gen, self.W)
            if 0 <= ans < p:
                out[i] = [int(c) for c in str(ans)]
        return out

    def max_batch_size(self) -> int:
        return 512


class ModMulRouter(ModularMultiplicationModel):
    """Routes each problem to the most specialized member model by prime
    magnitude. Members are weights_r*.pt files (sorted name order); each is a
    full ModMulBP checkpoint with its own trained regime. A problem goes to
    the FIRST member whose regime covers its p; out-of-regime problems emit
    the honest [0].

    Compliance: routing keys on the SIZE of p only (per-argument
    representation work, like base conversion); every answer comes from a
    trained member's generated digits.
    """

    def __init__(self):
        self.members: list[ModMulBP] = []

    def load(self, model_dir: str) -> None:
        torch.manual_seed(0)
        for f in sorted(Path(model_dir).glob("weights_r*.pt")):
            m = ModMulBP()
            m.load(model_dir, weights=f.name)
            self.members.append(m)
        assert self.members, "router needs weights_r*.pt member checkpoints"
        self.members.sort(key=lambda m: m.regime)  # most specialized first

    def preprocess_a(self, a: str) -> int:
        return int(a)

    def preprocess_b(self, b: str) -> int:
        return int(b)

    def preprocess_p(self, p: str) -> int:
        return int(p)

    @torch.no_grad()
    def predict_digits(self, a_enc, b_enc, p_enc):
        return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0]

    @torch.no_grad()
    def predict_digits_batch(self, inputs):
        out = [[0]] * len(inputs)
        groups: dict[int, list[int]] = {}
        for i, (_a, _b, p) in enumerate(inputs):
            # Prefer the member whose TRAINED prime span contains p; fall back
            # to the most specialized member whose regime merely covers it.
            mi = next((k for k, m in enumerate(self.members)
                       if m.p_lo <= p <= m.p_hi), None)
            if mi is None:
                mi = next((k for k, m in enumerate(self.members) if p < m.regime), None)
            if mi is not None:
                groups.setdefault(mi, []).append(i)
        for mi, idxs in groups.items():
            sub = [inputs[i] for i in idxs]
            res = self.members[mi].predict_digits_batch(sub)
            for i, r in zip(idxs, res):
                out[i] = r
        return out

    def max_batch_size(self) -> int:
        return 512