| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| class SequenceRegressor(nn.Module): | |
| def __init__(self, cell: str) -> None: | |
| super().__init__() | |
| self.cell = cell | |
| if cell == "rnn": | |
| hidden = 67 | |
| self.recurrent = nn.RNN( | |
| 2, hidden, nonlinearity="tanh", batch_first=True | |
| ) | |
| elif cell == "lstm": | |
| hidden = 32 | |
| self.recurrent = nn.LSTM(2, hidden, batch_first=True) | |
| elif cell == "gru": | |
| hidden = 36 | |
| self.recurrent = nn.GRU(2, hidden, batch_first=True) | |
| else: | |
| raise ValueError(f"Unknown recurrent cell: {cell}") | |
| self.readout = nn.Linear(hidden, 1) | |
| def forward(self, sequence: torch.Tensor) -> torch.Tensor: | |
| hidden, _ = self.recurrent(sequence) | |
| return self.readout(hidden[:, -1]).squeeze(-1) | |
| def parameter_count(module: nn.Module) -> int: | |
| return sum(parameter.numel() for parameter in module.parameters()) | |