From 72710c21d909c3b0573ea9ba8b7236683516342f Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Tue, 28 Jan 2020 10:10:47 +0100 Subject: [PATCH] api to set the transformer dim and ff_dim --- pts/model/transformer/transformer_estimator.py | 12 ++++++++---- pts/model/transformer/transformer_network.py | 9 +++++---- 2 files changed, 13 insertions(+), 8 deletions(-) diff --git a/pts/model/transformer/transformer_estimator.py b/pts/model/transformer/transformer_estimator.py index a2ca055..6583bb7 100644 --- a/pts/model/transformer/transformer_estimator.py +++ b/pts/model/transformer/transformer_estimator.py @@ -49,7 +49,8 @@ class TransformerEstimator(PTSEstimator): cardinality: Optional[List[int]] = None, embedding_dimension: int = 20, distr_output: DistributionOutput = StudentTOutput(), - dim_feedforward: int = 128, + d_model: int = 32, + dim_feedforward_scale: int = 4, act_type: str = "gelu", num_heads: int = 8, num_encoder_layers: int = 3, @@ -90,9 +91,10 @@ class TransformerEstimator(PTSEstimator): self.history_length = self.context_length + max(self.lags_seq) self.scaling = scaling + self.d_model = d_model self.num_heads = num_heads self.act_type = act_type - self.dim_feedforward = dim_feedforward + self.dim_feedforward_scale = dim_feedforward_scale self.num_encoder_layers = num_encoder_layers self.num_decoder_layers = num_decoder_layers @@ -175,7 +177,8 @@ class TransformerEstimator(PTSEstimator): num_heads=self.num_heads, act_type=self.act_type, dropout_rate=self.dropout_rate, - dim_feedforward=self.dim_feedforward, + d_model=self.d_model, + dim_feedforward_scale=self.dim_feedforward_scale, num_encoder_layers=self.num_encoder_layers, num_decoder_layers=self.num_decoder_layers, history_length=self.history_length, @@ -199,7 +202,8 @@ class TransformerEstimator(PTSEstimator): num_heads=self.num_heads, act_type=self.act_type, dropout_rate=self.dropout_rate, - dim_feedforward=self.dim_feedforward, + d_model=self.d_model, + dim_feedforward_scale=self.dim_feedforward_scale, num_encoder_layers=self.num_encoder_layers, num_decoder_layers=self.num_decoder_layers, history_length=self.history_length, diff --git a/pts/model/transformer/transformer_network.py b/pts/model/transformer/transformer_network.py index cb1c5e1..a8b44ba 100644 --- a/pts/model/transformer/transformer_network.py +++ b/pts/model/transformer/transformer_network.py @@ -21,10 +21,11 @@ class TransformerNetwork(nn.Module): def __init__( self, input_size: int, + d_model: int, num_heads: int, act_type: str, dropout_rate: float, - dim_feedforward: int, + dim_feedforward_scale: int, num_encoder_layers: int, num_decoder_layers: int, history_length: int, @@ -57,11 +58,11 @@ class TransformerNetwork(nn.Module): self.target_shape = distr_output.event_shape self.transformer = nn.Transformer( - d_model=input_size, + d_model=d_model, nhead=num_heads, num_encoder_layers=num_encoder_layers, num_decoder_layers=num_decoder_layers, - dim_feedforward=dim_feedforward, + dim_feedforward=dim_feedforward_scale*d_model, dropout=dropout_rate, activation=act_type, ) @@ -275,7 +276,7 @@ class TransformerTrainingNetwork(TransformerNetwork): dec_output = self.transformer.decoder( dec_input, enc_out, # memory - self.upper_triangular_mask(self.prediction_length), # mask + memory_mask=self.upper_triangular_mask(self.prediction_length), # target mask or memory mask? ) # compute loss