"""MiniNeuralHorner development model for modular multiplication. The learned component is one modulus-conditioned recurrent transition cell. It predicts the next binary residue state for s_next = (2 * s + d * x) mod p. A fixed Horner schedule applies that cell to reduce both operands and then multiply the two residues. The emitted answer is a list of base-2 digits. The SAIR evaluation harness performs the final digit decoding. """ from __future__ import annotations import hashlib import json from pathlib import Path import torch from torch import nn from modchallenge.interface.base_model import ModularMultiplicationModel _MASK32 = (1 << 32) - 1 _CHECKPOINT_SCHEMA = "mini-neuralhorner-inference-v1" _EXPECTED_CONFIG = { "bidirectional": True, "dmodel": 96, "hidden": 61, "num_layers": 2, } _EXPECTED_PARAMETERS = 126_603 _EXPECTED_WIDTH = 2_048 _EXPECTED_TENSOR_SHA256 = ( "7d1768ae1260f750e0a80ec93d98f86a80d441e479ce29e0a8e21fd098c742a3" ) _QUALIFIED_CHECKPOINT_SHA256 = ( "d296b711bb6a7faaa1dd81e05478cfa75f11071c42a8c36fbf60e758ee7eb407" ) def _to_bits_small(values: torch.Tensor, width: int) -> torch.Tensor: shifts = torch.arange(width - 1, -1, -1, device=values.device) return (values[:, None] >> shifts[None, :]) & 1 def to_bits_limbs(values: list[int], device: torch.device, width: int) -> torch.Tensor: """Convert nonnegative Python integers to MSB-first bits without int64 overflow.""" limb_count = (width + 31) // 32 columns = [] for limb_index in range(limb_count - 1, -1, -1): limb = torch.tensor( [(value >> (32 * limb_index)) & _MASK32 for value in values], dtype=torch.int64, device=device, ) columns.append(_to_bits_small(limb, 32)) bits = torch.cat(columns, dim=1) excess = limb_count * 32 - width return bits[:, excess:] if excess else bits def _tensor_digest(state_dict: dict[str, torch.Tensor]) -> str: """Hash tensor names, dtypes, shapes, and raw values in a stable order.""" digest = hashlib.sha256() for name in sorted(state_dict): tensor = state_dict[name].detach().cpu().contiguous() header = json.dumps( {"dtype": str(tensor.dtype), "name": name, "shape": list(tensor.shape)}, sort_keys=True, separators=(",", ":"), ).encode("utf-8") raw = tensor.numpy().tobytes(order="C") digest.update(len(header).to_bytes(8, "big")) digest.update(header) digest.update(len(raw).to_bytes(8, "big")) digest.update(raw) return digest.hexdigest() class TransitionCell(nn.Module): def __init__( self, dmodel: int, hidden: int, num_layers: int, bidirectional: bool, ) -> None: super().__init__() directions = 2 if bidirectional else 1 self.in_proj = nn.Linear(3, dmodel) self.d_emb = nn.Embedding(2, dmodel) self.gru = nn.GRU( dmodel, hidden, num_layers=num_layers, batch_first=True, bidirectional=bidirectional, ) self.head = nn.Linear(directions * hidden, 1) def forward( self, features: torch.Tensor, control: torch.Tensor, ) -> torch.Tensor: embedded = self.in_proj(features) + self.d_emb(control)[:, None, :] hidden, _ = self.gru(embedded) return self.head(hidden).squeeze(-1) def _bits_of(value: int) -> list[int]: if value <= 0: return [0] digits = [] while value: digits.append(value & 1) value >>= 1 digits.reverse() return digits class MiniNeuralHorner(ModularMultiplicationModel): """SAIR interface adapter for the 126,603-parameter transition cell.""" def __init__(self) -> None: self.model: TransitionCell | None = None self.device = torch.device("cpu") self.width = _EXPECTED_WIDTH self._sequence_width = 32 def load(self, model_dir: str) -> None: checkpoint_path = Path(model_dir) / "weights.pt" checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True) if checkpoint.get("schema") != _CHECKPOINT_SCHEMA: raise ValueError("unsupported MiniNeuralHorner checkpoint schema") if checkpoint.get("config") != _EXPECTED_CONFIG: raise ValueError("checkpoint architecture does not match the packaged model") if checkpoint.get("L") != _EXPECTED_WIDTH: raise ValueError("checkpoint inference width does not match the packaged model") provenance = checkpoint.get("provenance", {}) if provenance.get("qualified_checkpoint_sha256") != _QUALIFIED_CHECKPOINT_SHA256: raise ValueError("checkpoint provenance does not match the qualified source") state_dict = checkpoint.get("state_dict") if not isinstance(state_dict, dict): raise ValueError("checkpoint is missing its state_dict") parameter_count = sum(tensor.numel() for tensor in state_dict.values()) if parameter_count != _EXPECTED_PARAMETERS: raise ValueError("checkpoint parameter count is invalid") if _tensor_digest(state_dict) != _EXPECTED_TENSOR_SHA256: raise ValueError("checkpoint tensor digest is invalid") if torch.cuda.is_available(): self.device = torch.device("cuda") torch.backends.cudnn.benchmark = False torch.backends.cudnn.deterministic = True torch.backends.cudnn.allow_tf32 = False torch.backends.cuda.matmul.allow_tf32 = False elif torch.backends.mps.is_available(): self.device = torch.device("mps") self.model = TransitionCell(**_EXPECTED_CONFIG) self.model.load_state_dict(state_dict, strict=True) self.model.to(self.device) self.model.eval() def preprocess_a(self, a: str) -> list[int]: return _bits_of(int(a)) def preprocess_b(self, b: str) -> list[int]: return _bits_of(int(b)) def preprocess_p(self, p: str) -> int: return int(p) @torch.no_grad() def predict_digits( self, a_enc: list[int], b_enc: list[int], p_enc: int, ) -> list[int]: return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0] @torch.no_grad() def predict_digits_batch( self, inputs: list[tuple[list[int], list[int], int]], ) -> list[list[int]]: if self.model is None: raise RuntimeError("load() must be called before inference") max_operand_bits = 4 * self.width outputs: list[list[int]] = [[0] for _ in inputs] valid_indices = [] a_bit_lists = [] b_bit_lists = [] moduli = [] for index, (a_enc, b_enc, p_enc) in enumerate(inputs): modulus = int(p_enc) a_bits = list(a_enc) b_bits = list(b_enc) if ( modulus < 2 or modulus >= (1 << self.width) or len(a_bits) > max_operand_bits or len(b_bits) > max_operand_bits ): continue valid_indices.append(index) a_bit_lists.append(a_bits) b_bit_lists.append(b_bits) moduli.append(modulus) if not valid_indices: return outputs maximum_modulus_bits = max(modulus.bit_length() for modulus in moduli) self._sequence_width = min( self.width, max(32, ((maximum_modulus_bits + 31) // 32) * 32), ) modulus_bits = to_bits_limbs( moduli, self.device, self._sequence_width, ).float() a_residues = self._reduce(a_bit_lists, modulus_bits) b_residues = self._reduce(b_bit_lists, modulus_bits) product = self._multiply(a_residues, b_residues, modulus_bits) for result_index, input_index in enumerate(valid_indices): outputs[input_index] = [int(bit) for bit in product[result_index].long().tolist()] return outputs def max_batch_size(self) -> int: return 256 def _step( self, state_bits: torch.Tensor, multiplicand_bits: torch.Tensor, modulus_bits: torch.Tensor, control: torch.Tensor, ) -> torch.Tensor: if self.model is None: raise RuntimeError("model is not loaded") features = torch.stack( [state_bits, multiplicand_bits, modulus_bits], dim=-1, ) logits = self.model(features, control) return (torch.sigmoid(logits) > 0.5).float() def _reduce( self, bit_lists: list[list[int]], modulus_bits: torch.Tensor, ) -> torch.Tensor: batch_size = len(bit_lists) operand_width = max(len(bits) for bits in bit_lists) padded = torch.zeros( (batch_size, operand_width), dtype=torch.long, device=self.device, ) for row, bits in enumerate(bit_lists): if bits: padded[row, operand_width - len(bits) :] = torch.tensor( bits, dtype=torch.long, device=self.device, ) state_bits = torch.zeros( (batch_size, self._sequence_width), device=self.device, ) one_bits = to_bits_limbs( [1] * batch_size, self.device, self._sequence_width, ).float() for position in range(operand_width): state_bits = self._step( state_bits, one_bits, modulus_bits, padded[:, position], ) return state_bits def _multiply( self, a_residue: torch.Tensor, b_residue: torch.Tensor, modulus_bits: torch.Tensor, ) -> torch.Tensor: batch_size = a_residue.shape[0] state_bits = torch.zeros( (batch_size, self._sequence_width), device=self.device, ) b_digits = b_residue.long() for position in range(self._sequence_width): state_bits = self._step( state_bits, a_residue, modulus_bits, b_digits[:, position], ) return state_bits