from __future__ import annotations from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from modchallenge.interface.base_model import ModularMultiplicationModel MAX_P_BITS = 32 PAD_HEAD = 3 REDUCE_FEATURES = 6 ADD_FEATURES = 5 def _bits_of(n: int) -> list[int]: return [int(c) for c in bin(n)[2:]] def linear_scan(alpha: torch.Tensor, beta: torch.Tensor) -> torch.Tensor: A, B = alpha, beta n = A.shape[1] off = 1 while off < n: a_prev = F.pad(A, (0, 0, off, 0), value=1.0)[:, :n] b_prev = F.pad(B, (0, 0, off, 0), value=0.0)[:, :n] B = A * b_prev + B A = A * a_prev off <<= 1 return B class GateFn(nn.Module): def __init__(self, mode: str = "hard"): super().__init__() self.mode = mode def forward(self, z: torch.Tensor) -> torch.Tensor: if self.mode == "soft": return torch.sigmoid(z) hard = (z > 0).to(z.dtype) if self.mode == "ste" and self.training: soft = torch.sigmoid(z) return hard + soft - soft.detach() return hard class BiScanBlock(nn.Module): def __init__(self, d_model: int, d_scan: int, gate: GateFn): super().__init__() self.gate = gate self.proj_f = nn.Linear(d_model, 2 * d_scan) self.proj_b = nn.Linear(d_model, 2 * d_scan) self.out = nn.Linear(2 * d_scan, d_model) self.mlp = nn.Sequential( nn.Linear(d_model, 2 * d_model), nn.ReLU(), nn.Linear(2 * d_model, d_model), ) self.scan_noise = 0.0 def forward(self, u: torch.Tensor) -> torch.Tensor: zf = self.proj_f(u) zb = self.proj_b(u) af, bf = zf.chunk(2, dim=-1) ab, bb = zb.chunk(2, dim=-1) hf = linear_scan(self.gate(af), bf) hb = linear_scan(self.gate(ab.flip(1)), bb.flip(1)).flip(1) h = torch.cat([hf, hb], dim=-1) if self.training: self.last_h_l1 = h.abs().mean() if self.training and self.scan_noise > 0: h = h + torch.randn_like(h) * self.scan_noise u = u + self.out(h) u = u + self.mlp(u) return u class BitCell(nn.Module): def __init__(self, n_features: int, n_borrow: int, n_q: int, d_model: int = 32, d_scan: int = 16, n_blocks: int = 3, gate_mode: str = "hard"): super().__init__() self.gate = GateFn(gate_mode) self.embed = nn.Linear(n_features, d_model) self.pre_mlp = nn.Sequential( nn.Linear(d_model, 2 * d_model), nn.ReLU(), nn.Linear(2 * d_model, d_model), ) self.blocks = nn.ModuleList( BiScanBlock(d_model, d_scan, self.gate) for _ in range(n_blocks) ) self.head = nn.Linear(d_model, 1) self.head_carry = nn.Linear(d_model, 1) self.head_sum = nn.Linear(d_model, 1) self.head_borrow = nn.Linear(d_model, n_borrow) self.head_q = nn.Linear(d_model, n_q) self.config = dict(n_features=n_features, n_borrow=n_borrow, n_q=n_q, d_model=d_model, d_scan=d_scan, n_blocks=n_blocks) def trunk(self, feats: torch.Tensor): u = self.embed(feats) u = u + self.pre_mlp(u) taps = [] for blk in self.blocks: u = blk(u) taps.append(u) return u, taps def forward(self, feats: torch.Tensor) -> torch.Tensor: u, _ = self.trunk(feats) return self.head(u).squeeze(-1) def forward_train(self, feats: torch.Tensor): u, taps = self.trunk(feats) return { "bits": self.head(u).squeeze(-1), "carry": self.head_carry(taps[0]).squeeze(-1), "sum": self.head_sum(taps[0]).squeeze(-1), "borrow": self.head_borrow(taps[1]), "q": self.head_q(taps[-1].mean(dim=1)), } def make_reduce_cell(gate_mode: str = "hard", **kw) -> BitCell: return BitCell(REDUCE_FEATURES, n_borrow=3, n_q=4, gate_mode=gate_mode, **kw) def make_add_cell(gate_mode: str = "hard", **kw) -> BitCell: kw.setdefault("n_blocks", 2) return BitCell(ADD_FEATURES, n_borrow=1, n_q=2, gate_mode=gate_mode, **kw) def shift_bits(t: torch.Tensor, k: int) -> torch.Tensor: if k == 0: return t return torch.cat([t[:, k:], t.new_zeros(t.shape[0], k)], dim=1) def _flags(B: int, N: int, dev) -> tuple[torch.Tensor, torch.Tensor]: is_msb = torch.zeros(B, N, device=dev) is_msb[:, 0] = 1.0 is_lsb = torch.zeros(B, N, device=dev) is_lsb[:, -1] = 1.0 return is_msb, is_lsb def reduce_features(x, p, p3): B, N = x.shape is_msb, is_lsb = _flags(B, N, x.device) return torch.stack([x, p, shift_bits(p, 1), p3, is_msb, is_lsb], dim=-1) def add_features(x, y, g): B, N = x.shape is_msb, is_lsb = _flags(B, N, x.device) return torch.stack([x, y, g.unsqueeze(1).expand(B, N), is_msb, is_lsb], dim=-1) class BitStreamMachine: def __init__(self, reduce_cell: BitCell, add_cell: BitCell, device: torch.device): self.reduce_cell = reduce_cell self.add_cell = add_cell self.device = device @torch.no_grad() def _rstep(self, x, p, p3): logits = self.reduce_cell(reduce_features(x, p, p3)) return (logits > 0).to(x.dtype) @torch.no_grad() def _astep(self, x, y, g): logits = self.add_cell(add_features(x, y, g)) return (logits > 0).to(x.dtype) @torch.no_grad() def run(self, a_bits, b_bits, p_bits, p3_bits): B, N = p_bits.shape L = a_bits.shape[1] ops = torch.cat([a_bits, b_bits], dim=0) p2r = torch.cat([p_bits, p_bits], dim=0) p32r = torch.cat([p3_bits, p3_bits], dim=0) X = p2r.new_zeros(2 * B, N) for t in range(0, L, 2): x = torch.cat([X[:, 2:], ops[:, t: t + 2]], dim=1) X = self._rstep(x, p2r, p32r) ra, rb = X[:B], X[B:] Z = p_bits.new_zeros(B, N) for t in range(PAD_HEAD, N): s = self._astep(shift_bits(Z, 1), rb, ra[:, t]) Z = self._rstep(s, p_bits, p3_bits) return Z class BitStreamModel(ModularMultiplicationModel): def __init__(self): self.machine: BitStreamMachine | None = None self.device = torch.device("cpu") def load(self, model_dir: str) -> None: torch.set_grad_enabled(False) ckpt = torch.load( Path(model_dir) / "weights.pt", map_location="cpu", weights_only=True, ) rcell = make_reduce_cell() rcell.load_state_dict(ckpt["reduce_state_dict"], strict=True) rcell.eval() acell = make_add_cell() acell.load_state_dict(ckpt["add_state_dict"], strict=True) acell.eval() self.machine = BitStreamMachine(rcell, acell, self.device) def preprocess_a(self, a: str): return _bits_of(int(a)) def preprocess_b(self, b: str): return _bits_of(int(b)) def preprocess_p(self, p: str): v = int(p) return {"p": _bits_of(v), "p3": _bits_of((v << 1) + v)} @torch.no_grad() def predict_digits(self, a_enc, b_enc, p_enc) -> list[int]: return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0] @torch.no_grad() def predict_digits_batch(self, inputs) -> list[list[int]]: out: list[list[int] | None] = [None] * len(inputs) idx = [i for i, (_, _, pe) in enumerate(inputs) if len(pe["p"]) <= MAX_P_BITS] keep = set(idx) for i in range(len(inputs)): if i not in keep: out[i] = [0] if not idx: return [o if o is not None else [0] for o in out] sub = [inputs[i] for i in idx] n_p = max(len(pe["p"]) for _, _, pe in sub) + PAD_HEAD L = max(2, max(max(len(ae), len(be)) for ae, be, _ in sub)) L += L % 2 def pack(rows: list[list[int]], width: int) -> torch.Tensor: t = torch.zeros(len(rows), width) for r, bits in enumerate(rows): if bits: t[r, width - len(bits):] = torch.tensor( bits, dtype=torch.float32) return t a_t = pack([ae for ae, _, _ in sub], L) b_t = pack([be for _, be, _ in sub], L) p_t = pack([pe["p"] for _, _, pe in sub], n_p) p3_t = pack([pe["p3"] for _, _, pe in sub], n_p) z = self.machine.run(a_t, b_t, p_t, p3_t) z_int = z.to(torch.int64).tolist() for row, i in enumerate(idx): out[i] = [int(v) for v in z_int[row]] return [o if o is not None else [0] for o in out] def max_batch_size(self) -> int: return 128