"""Weight-tied RNN: ONE 2-layer cell reused for input read-in, K compute steps, and output decoding. Spec pseudocode (design/experiment-spec.md): state = embed(input_number) for step in range(max_steps): state = step_module(state) # same weights, every iteration if halt_condition(state): break output = project(state) Everything recurrent is the SAME cell (per design review: no un-tied GRU decoder, no mean-pool blur — order reaches the cell via sinusoidal position added per digit). ACT halting applies only to the K compute steps. """ import torch import torch.nn as nn import torch.nn.functional as F from src.config import Config from src.model_api import PrimeModel class TiedRNN(PrimeModel): def __init__(self, cfg: Config): super().__init__() self.cfg = cfg d = cfg.d_model self.embed = nn.Embedding(cfg.vocab + 1, d) # +1 row = pad self.cell_ln = nn.LayerNorm(d) self.cell_w1 = nn.Linear(d, d) self.cell_w2 = nn.Linear(d, d) self.ln1 = nn.LayerNorm(d) self.halt_head = nn.Linear(d, 1) self.out_head = nn.Linear(d, cfg.vocab) def _cell_step(self, h: torch.Tensor) -> torch.Tensor: return h + self.cell_w2(F.gelu(self.cell_w1(self.cell_ln(h)))) @staticmethod def _sinusoidal(T: int, d: int) -> torch.Tensor: pe = torch.zeros(T, d) pos = torch.arange(T).float().unsqueeze(1) i = torch.arange(d).float().unsqueeze(0) pe[:, 0::2] = torch.sin(pos / 10000 ** (2 * i[:, 0::2] / d)) pe[:, 1::2] = torch.cos(pos / 10000 ** (2 * i[:, 1::2] / d)) return pe def _encode(self, x: torch.Tensor) -> torch.Tensor: """Read input digits through the tied cell. Pad positions are exact no-ops (state update masked) so batch composition cannot change an example's state.""" B, T = x.shape d = self.cfg.d_model pos = self._sinusoidal(T, d).to(x.device) # (T,d) e = self.embed(x) # (B,T,d) mask = (x != self.cfg.pad_id).float().unsqueeze(-1) # (B,T,1) h = torch.zeros(B, d, device=x.device) for t in range(T): h_new = self._cell_step(h + (e[:, t] + pos[t]) * mask[:, t]) h = mask[:, t] * h_new + (1 - mask[:, t]) * h # pad step = no-op return h def _run_compute(self, h0: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """K tied compute steps with ACT learned halting. Returns (final_state, mean_steps).""" cfg = self.cfg B = h0.shape[0] device = h0.device if not cfg.halting: h = h0 for _ in range(cfg.max_steps): h = self._cell_step(h) steps = h0.new_full((B,), float(cfg.max_steps)) return h, steps h_list, p_list = [], [] h = h0 for t in range(cfg.max_steps): h = self._cell_step(h) p = torch.sigmoid(self.halt_head(self.ln1(h))).squeeze(-1) # (B,) if t < cfg.min_steps - 1: p = p * 0.0 h_list.append(h) p_list.append(p) # ACT aggregation (Graves 2016): w_t = p_t * prod_{s dict: cfg = self.cfg h = self._encode(x) h, steps = self._run_compute(h) # decode output digits through the SAME tied cell (teacher-forced during training) e_out = self.embed(y_in) # (B,T_out,d) outs = [] for t in range(y_in.shape[1]): h = self._cell_step(h + e_out[:, t]) outs.append(self.out_head(h)) logits = torch.stack(outs, dim=1) # (B,T_out,vocab) return {"logits": logits, "halt_steps": steps}