summaryrefslogtreecommitdiff
path: root/src/model_api.py
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-08-14 13:11:02 +0100
committerVoid Agent <void@jayrup.hermes>2026-08-14 13:11:02 +0100
commit898a0570dfe619ec2bcf330b07dce9b23cb63d54 (patch)
tree423468249add4570856dc20240d2e30bf60cdcb0 /src/model_api.py
parent6e7268b66b407ea3603fc9128805d132826a769f (diff)
implement Experiment 1: data pipeline, tied-RNN w/ ACT halting, transformer baseline, train/eval/plot, 16 tests
Diffstat (limited to 'src/model_api.py')
-rw-r--r--src/model_api.py38
1 files changed, 38 insertions, 0 deletions
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:]