diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config.py | 3 | ||||
| -rw-r--r-- | src/eval.py | 4 | ||||
| -rw-r--r-- | src/models/rnn.py | 54 | ||||
| -rw-r--r-- | src/train.py | 9 |
4 files changed, 42 insertions, 28 deletions
diff --git a/src/config.py b/src/config.py index 74cb4d3..c0532bc 100644 --- a/src/config.py +++ b/src/config.py @@ -19,7 +19,8 @@ class Config: halting: bool = True # False -> fixed-K ablation halt_eps: float = 0.05 # ACT cumulative threshold (documented; K small so no early break) halt_penalty: float = 0.01 # lambda on mean steps - halt_warmup_steps: int = 1000 # penalty off before this (anti-collapse) + halt_warmup_steps: int = 1000 # penalty = 0 before this (anti-collapse) + halt_ramp_end_steps: int = 5000 # penalty ramps linearly 0 -> halt_penalty between warmup and here min_steps: int = 2 # ACT: halt prob forced to 0 before this step n_layers: int = 2 n_heads: int = 4 diff --git a/src/eval.py b/src/eval.py index 9d55437..4485523 100644 --- a/src/eval.py +++ b/src/eval.py @@ -91,8 +91,8 @@ def halting_report(model, cfg: Config) -> dict: gaps, steps = [], [] for n in range(cfg.range_start, cfg.range_end + 1): x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0) - h = model._initial_state(x) - _, s = model._run_cell(h) + h = model._encode(x) + _, s = model._run_compute(h) gaps.append(next_prime(n, primes) - n) steps.append(float(s.mean())) mean = float(np.mean(steps)) diff --git a/src/models/rnn.py b/src/models/rnn.py index 9947c66..19123b7 100644 --- a/src/models/rnn.py +++ b/src/models/rnn.py @@ -1,14 +1,15 @@ -"""Weight-tied RNN: one 2-layer cell applied K times, ACT learned halting, GRU digit decoder. +"""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 + state = step_module(state) # same weights, every iteration if halt_condition(state): break output = project(state) -Initial state: masked mean-pool of (digit embedding + sinusoidal positional encoding) -passed through a small MLP, so digit ORDER reaches the tied cell. +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 @@ -24,15 +25,16 @@ class TiedRNN(PrimeModel): self.cfg = cfg d = cfg.d_model self.embed = nn.Embedding(cfg.vocab + 1, d) # +1 row = pad - self.in_proj = nn.Sequential(nn.Linear(d, d), nn.GELU(), nn.Linear(d, d)) - self.ln1 = nn.LayerNorm(d) 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.decoder = nn.GRUCell(d, d) 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) @@ -42,35 +44,39 @@ class TiedRNN(PrimeModel): pe[:, 1::2] = torch.cos(pos / 10000 ** (2 * i[:, 1::2] / d)) return pe - def _initial_state(self, x: torch.Tensor) -> torch.Tensor: + def _encode(self, x: torch.Tensor) -> torch.Tensor: + """Read input digits through the tied cell (pad positions inject nothing).""" B, T = x.shape - mask = (x != self.cfg.pad_id).float().unsqueeze(-1) # (B,T,1) - pos = self._sinusoidal(T, self.cfg.d_model).to(x.device) # (T,d) - e = self.embed(x) + pos.unsqueeze(0) # (B,T,d) - h = (e * mask).sum(1) / mask.sum(1).clamp(min=1) # (B,d) - return self.in_proj(h) + 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 = self._cell_step(h + (e[:, t] + pos[t]) * mask[:, t]) + return h - def _run_cell(self, h0: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """Tied cell x K steps. Returns (final_state (B,d), mean_steps (B,)).""" + 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 = h + self.cell_w2(F.gelu(self.cell_w1(self.cell_ln(h)))) + h = self._cell_step(h) steps = h0.new_full((B,), float(cfg.max_steps)) return h, steps - # ACT: run all K steps, accumulate weighted average (K=20 -> no early break needed) h_list, p_list = [], [] h = h0 for t in range(cfg.max_steps): - h = h + self.cell_w2(F.gelu(self.cell_w1(self.cell_ln(h)))) + h = self._cell_step(h) p = torch.sigmoid(self.halt_head(self.ln1(h))).squeeze(-1) # (B,) if t < cfg.min_steps: p = p * 0.0 h_list.append(h) p_list.append(p) + # ACT aggregation (Graves 2016): w_t = p_t * prod_{s<t}(1-p_s); weights + remainder = 1 final = torch.zeros_like(h0) steps = torch.zeros(B, device=device) remaining = torch.ones(B, device=device) @@ -86,13 +92,13 @@ class TiedRNN(PrimeModel): def forward(self, x: torch.Tensor, y_in: torch.Tensor) -> dict: cfg = self.cfg - h = self._initial_state(x) - h, steps = self._run_cell(h) - # autoregressive digit decoder, teacher-forced during training - e = self.embed(y_in) # (B,T_out,d) + 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.decoder(e[:, t], h) + h = self._cell_step(h + e_out[:, t]) outs.append(self.out_head(h)) - logits = torch.stack(outs, dim=1) # (B,T_out,vocab) + logits = torch.stack(outs, dim=1) # (B,T_out,vocab) return {"logits": logits, "halt_steps": steps} diff --git a/src/train.py b/src/train.py index f0c63e9..975598d 100644 --- a/src/train.py +++ b/src/train.py @@ -133,7 +133,14 @@ def main() -> None: loss_tokens = ce(logits.reshape(-1, cfg.vocab), y_safe.reshape(-1)).reshape( logits.shape[0], -1) * batch["y_mask"].float() loss_tokens = loss_tokens.sum() / batch["y_mask"].sum().clamp(min=1) - lam = cfg.halt_penalty if (cfg.halting and step >= cfg.halt_warmup_steps) else 0.0 + if cfg.halting and step >= cfg.halt_warmup_steps: + if step >= cfg.halt_ramp_end_steps: + lam = cfg.halt_penalty + else: + frac = (step - cfg.halt_warmup_steps) / max(1, cfg.halt_ramp_end_steps - cfg.halt_warmup_steps) + lam = cfg.halt_penalty * frac + else: + lam = 0.0 halt = out["halt_steps"] penalty = halt.float().mean() * lam if halt is not None and lam > 0 else 0.0 loss = loss_tokens + penalty |
