Fix bug with device change
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user