fix typos

This commit is contained in:
Dr. Kashif Rasul
2020-01-28 10:30:53 +01:00
parent 32db813e41
commit 969b8c41b3
2 changed files with 5 additions and 5 deletions
@@ -47,7 +47,7 @@ class TransformerEstimator(PTSEstimator):
trainer: Trainer = Trainer(),
dropout_rate: float = 0.1,
cardinality: Optional[List[int]] = None,
embedding_dimension: int = 20,
embedding_dimension: List[int] = [20],
distr_output: DistributionOutput = StudentTOutput(),
d_model: int = 32,
dim_feedforward_scale: int = 4,
+4 -4
View File
@@ -33,7 +33,7 @@ class TransformerNetwork(nn.Module):
prediction_length: int,
distr_output: DistributionOutput,
cardinality: List[int],
embedding_dimension: int,
embedding_dimension: List[int],
lags_seq: List[int],
scaling: bool = True,
**kwargs,
@@ -55,8 +55,8 @@ class TransformerNetwork(nn.Module):
self.target_shape = distr_output.event_shape
self.encoder_input = nn.Dense(input_size, d_model)
self.decoder_input = nn.Dense(input_size, d_model)
self.encoder_input = nn.Linear(input_size, d_model)
self.decoder_input = nn.Linear(input_size, d_model)
self.transformer = nn.Transformer(
d_model=d_model,
@@ -281,7 +281,7 @@ class TransformerTrainingNetwork(TransformerNetwork):
dec_output = self.transformer.decoder(
self.decoder_input(dec_input),
enc_out, # memory
memory_mask=self.upper_triangular_mask(
tgt_mask=self.upper_triangular_mask(
self.prediction_length
), # target mask or memory mask?
)