diff --git a/Code/python/src/aiRNN/models.py b/Code/python/src/aiRNN/models.py index 859e06a..a63db4b 100644 --- a/Code/python/src/aiRNN/models.py +++ b/Code/python/src/aiRNN/models.py @@ -36,6 +36,14 @@ class BaseRNN(nn.Module): batch_first=True, device=device ) + elif rnn_type == "GRU": + self.rnn = nn.GRU( + input_size=hidden_size*3, + hidden_size=rnn_size, + num_layers=3, + batch_first=True, + device=device + ) else: raise ValueError("Unsupported rnn_type") @@ -132,4 +140,27 @@ class ThreeInputLSTM(BaseRNN): Predictions can be compared to targets with a specified `offset`. """ def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, device="cpu"): - super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="LSTM", device=device) \ No newline at end of file + super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="LSTM", device=device) + +class TwoInputGRU(BaseRNN): + """ + Many-to-many GRU with an initialization phase that consumes both: + - time inputs (always present) + - feature inputs (only during init_steps) + After init_steps, only time inputs are provided and the model predicts a sequence. + Predictions can be compared to targets with a specified `offset`. + """ + def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, device="cpu"): + super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="GRU", device=device) + +class ThreeInputGRU(BaseRNN): + """ + Many-to-many GRU with an initialization phase that consumes three inputs: + - time inputs (always present) + - feature inputs (only during init_steps) + - context inputs (always present + constant over time) + After init_steps, only time and context inputs are provided and the model predicts a sequence. + Predictions can be compared to targets with a specified `offset`. + """ + def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, device="cpu"): + super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="GRU", device=device) \ No newline at end of file