From 8fa0297ab7ddbf26c69de501d95602c01344ea9f Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Sat, 1 Feb 2020 18:04:01 +0100 Subject: [PATCH] scale input to network if flag is set --- pts/model/tempflow/tempflow_network.py | 5 ++++- .../transformer_tempflow/transformer_tempflow_network.py | 5 ++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/pts/model/tempflow/tempflow_network.py b/pts/model/tempflow/tempflow_network.py index 4a4fc31..acbe637 100644 --- a/pts/model/tempflow/tempflow_network.py +++ b/pts/model/tempflow/tempflow_network.py @@ -75,7 +75,10 @@ class TempFlowTrainingNetwork(nn.Module): num_embeddings=self.target_dim, embedding_dim=self.embed_dim ) - self.scaler = MeanScaler(keepdim=True) + if self.scaling: + self.scaler = MeanScaler(keepdim=True) + else: + self.scaler = NOPScaler(keepdim=True) @staticmethod def get_lagged_subsequences( diff --git a/pts/model/transformer_tempflow/transformer_tempflow_network.py b/pts/model/transformer_tempflow/transformer_tempflow_network.py index ee70607..e3340d7 100644 --- a/pts/model/transformer_tempflow/transformer_tempflow_network.py +++ b/pts/model/transformer_tempflow/transformer_tempflow_network.py @@ -82,7 +82,10 @@ class TransformerTempFlowTrainingNetwork(nn.Module): num_embeddings=self.target_dim, embedding_dim=self.embed_dim ) - self.scaler = MeanScaler(keepdim=True) + if self.scaling: + self.scaler = MeanScaler(keepdim=True) + else: + self.scaler = NOPScaler(keepdim=True) # mask self.register_buffer(