From 898a0570dfe619ec2bcf330b07dce9b23cb63d54 Mon Sep 17 00:00:00 2001 From: Void Agent Date: Fri, 14 Aug 2026 13:11:02 +0100 Subject: implement Experiment 1: data pipeline, tied-RNN w/ ACT halting, transformer baseline, train/eval/plot, 16 tests --- src/model_api.py | 38 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 38 insertions(+) create mode 100644 src/model_api.py (limited to 'src/model_api.py') diff --git a/src/model_api.py b/src/model_api.py new file mode 100644 index 0000000..66c2ace --- /dev/null +++ b/src/model_api.py @@ -0,0 +1,38 @@ +"""Model interface contract + build dispatch + shared greedy decode.""" +import torch +import torch.nn as nn + +from src.config import Config + + +class PrimeModel(nn.Module): + """Contract: forward(x, y_in) -> {"logits": (B,T_out,vocab), "halt_steps": (B,) or None}.""" + + def forward(self, x, y_in): + raise NotImplementedError + + def param_count(self) -> int: + return sum(p.numel() for p in self.parameters()) + + +def build_model(cfg: Config) -> PrimeModel: + if cfg.model == "rnn": + from src.models.rnn import TiedRNN + return TiedRNN(cfg) + if cfg.model == "transformer": + from src.models.transformer import TransformerBaseline + return TransformerBaseline(cfg) + raise ValueError(f"unknown model: {cfg.model}") + + +@torch.no_grad() +def greedy_decode(model: PrimeModel, x: torch.Tensor, cfg: Config, max_len: int | None = None) -> torch.Tensor: + """Autoregressive greedy decode of output digits. Returns (B, max_len) tokens (BOS stripped).""" + max_len = max_len or cfg.max_out_len + B = x.shape[0] + y_in = torch.full((B, 1), cfg.eos_id, dtype=torch.long, device=x.device) + for _ in range(max_len): + out = model(x, y_in) + nxt = out["logits"][:, -1].argmax(-1) + y_in = torch.cat([y_in, nxt[:, None]], dim=1) + return y_in[:, 1:] -- cgit v1.2.3