| """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)) |
| |
| |
| |
| |
| |
| |
| from train_arith_bp_supervised import TrainConfig, build_model |
| import encoding as enc |
| import composed_encoding as cenc |
| import kvgen |
|
|
| from modchallenge.interface.base_model import ModularMultiplicationModel |
|
|
|
|
| class ModMulBP(ModularMultiplicationModel): |
| def __init__(self): |
| self.model = None |
| self.W = 1 |
| self.device = None |
| self.regime = 10 |
| self.scratchpad = False |
| self.school = False |
|
|
| def load(self, model_dir: str, weights: str = "weights.pt") -> None: |
| torch.manual_seed(0) |
| 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: |
| |
| |
| |
| self.regime = 10 ** (2 * self.W) |
| self.p_lo, self.p_hi = 2, 2 ** 4096 |
| return |
| |
| |
| 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 |
|
|
| |
| 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( |
| x, y, p, W, 10, mulonly=True) |
| n_gen = cenc.gen_len(self.W, mulonly=True) |
| decode_fn = lambda g, W: cenc.decode_answer( |
| g, W, 10, mulonly=True) |
| elif self.composed: |
| prompt_fn = lambda p, x, y, W: cenc.prompt_str( |
| 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( |
| 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: |
| continue |
| elif p >= self.regime: |
| continue |
| prompts.append(prompt_fn(p, x, y, self.W)) |
| idx.append((i, x, y, p)) |
| if not prompts: |
| return out |
| |
| |
| ids = torch.tensor([list(s.encode("latin1")) for s in prompts], |
| dtype=torch.long, device=self.device) |
| plen = len(prompts[0]) |
| |
| |
| |
| |
| 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") |
| |
| |
| |
| 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) |
|
|
| 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): |
| |
| |
| 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 |
|
|