File size: 1,011 Bytes
ddaaeb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
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())