api to set the transformer dim and ff_dim

This commit is contained in:
Dr. Kashif Rasul
2020-01-28 10:10:47 +01:00
parent 8d92ae4b89
commit 72710c21d9
2 changed files with 13 additions and 8 deletions
@@ -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,
+5 -4
View File
@@ -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