diff --git a/Code/python/src/aiRNN/models.py b/Code/python/src/aiRNN/models.py index 640c164..326965c 100644 --- a/Code/python/src/aiRNN/models.py +++ b/Code/python/src/aiRNN/models.py @@ -4,26 +4,28 @@ import torch.nn as nn class BaseRNN(nn.Module): def __init__(self, time_in, feat_in, context_in, hidden_size, - rnn_size, out_size, rnn_type="RNN"): + rnn_size, out_size, rnn_type="RNN", device="cpu"): super().__init__() - self.time_proj = nn.Linear(time_in, hidden_size) - self.feat_proj = nn.Linear(feat_in, hidden_size) + self.time_proj = nn.Linear(time_in, hidden_size, device=device) + self.feat_proj = nn.Linear(feat_in, hidden_size, device=device) self.context_proj = ( - nn.Linear(context_in, hidden_size) if context_in is not None else None + nn.Linear(context_in, hidden_size, device=device) if context_in is not None else None ) if rnn_type == "RNN": self.rnn = nn.RNN( input_size=hidden_size, hidden_size=rnn_size, - batch_first=True + batch_first=True, + device=device ) elif rnn_type == "LSTM": self.rnn = nn.LSTM( input_size=hidden_size, hidden_size=rnn_size, - batch_first=True + batch_first=True, + device=device ) else: raise ValueError("Unsupported rnn_type") @@ -31,7 +33,7 @@ class BaseRNN(nn.Module): readout_in = rnn_size + time_in if context_in is not None: readout_in += context_in - self.readout = nn.Linear(readout_in, out_size) + self.readout = nn.Linear(readout_in, out_size, device=device) self.context_in = context_in self.rnn_size = rnn_size