From de28c4a880b53324a4208c7de54a0ec3f2829b80 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Fri, 21 Nov 2025 08:09:36 +0100 Subject: [PATCH] Fix bug with device change --- Code/python/src/aiRNN/models.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) 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