diff options
| author | Void Agent <void@jayrup.hermes> | 2026-08-14 13:11:02 +0100 |
|---|---|---|
| committer | Void Agent <void@jayrup.hermes> | 2026-08-14 13:11:02 +0100 |
| commit | 898a0570dfe619ec2bcf330b07dce9b23cb63d54 (patch) | |
| tree | 423468249add4570856dc20240d2e30bf60cdcb0 /src/model_api.py | |
| parent | 6e7268b66b407ea3603fc9128805d132826a769f (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.py | 38 |
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:] |
