summaryrefslogtreecommitdiff
path: root/src/model_api.py
diff options
context:
space:
mode:
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:]