| 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 |
|
|