diff --git a/Code/python/src/aiRNN/models.py b/Code/python/src/aiRNN/models.py index 326965c..39ab19d 100644 --- a/Code/python/src/aiRNN/models.py +++ b/Code/python/src/aiRNN/models.py @@ -91,8 +91,8 @@ class TwoInputRNN(BaseRNN): 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): - super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size) + 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, device=device) class ThreeInputRNN(BaseRNN): """ @@ -103,8 +103,8 @@ class ThreeInputRNN(BaseRNN): 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): - super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size) + 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, device=device) class TwoInputLSTM(BaseRNN): """ @@ -114,8 +114,8 @@ class TwoInputLSTM(BaseRNN): 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): - super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="LSTM") + 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="LSTM", device=device) class ThreeInputLSTM(BaseRNN): """ @@ -126,5 +126,5 @@ class ThreeInputLSTM(BaseRNN): 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): - 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") \ No newline at end of file + 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