Fix bug with device change

This commit is contained in:
2025-11-21 08:09:36 +01:00
parent 240ed75e3a
commit de28c4a880
+8 -8
View File
@@ -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")
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)