added static real features

This commit is contained in:
Kashif Rasul
2020-01-27 21:19:16 +01:00
parent 064baee01c
commit 8d92ae4b89
4 changed files with 24 additions and 2 deletions
+11 -2
View File
@@ -59,6 +59,7 @@ class TransformerEstimator(PTSEstimator):
time_features: Optional[List[TimeFeature]] = None,
use_feat_dynamic_real: bool = False,
use_feat_static_cat: bool = False,
use_feat_static_real: bool = False,
num_parallel_samples: int = 100,
) -> None:
super().__init__(trainer=trainer)
@@ -73,6 +74,7 @@ class TransformerEstimator(PTSEstimator):
self.dropout_rate = dropout_rate
self.use_feat_dynamic_real = use_feat_dynamic_real
self.use_feat_static_cat = use_feat_static_cat
self.use_feat_static_real = use_feat_static_real
self.cardinality = cardinality if use_feat_static_cat else [1]
self.embedding_dimension = embedding_dimension
self.num_parallel_samples = num_parallel_samples
@@ -97,21 +99,28 @@ class TransformerEstimator(PTSEstimator):
def create_transformation(self) -> Transformation:
remove_field_names = [
FieldName.FEAT_DYNAMIC_CAT,
FieldName.FEAT_STATIC_REAL,
]
if not self.use_feat_dynamic_real:
remove_field_names.append(FieldName.FEAT_DYNAMIC_REAL)
if not self.use_feat_static_real:
remove_field_names.append(FieldName.FEAT_STATIC_REAL)
return Chain(
[RemoveFields(field_names=remove_field_names)]
+ (
[SetField(output_field=FieldName.FEAT_STATIC_CAT, value=[0])]
if not self.use_feat_static_cat
else []
) + (
[SetField(output_field=FieldName.FEAT_STATIC_REAL, value=[0.0])]
if not self.use_feat_static_real
else []
)
+ [
AsNumpyArray(field=FieldName.FEAT_STATIC_CAT,
expected_ndim=1, dtype=np.long),
AsNumpyArray(
field=FieldName.FEAT_STATIC_REAL, expected_ndim=1, dtype=self.dtype,
),
AsNumpyArray(
field=FieldName.TARGET,
# in the following line, we add 1 for the time dimension
@@ -121,6 +121,7 @@ class TransformerNetwork(nn.Module):
def create_network_input(
self,
feat_static_cat: torch.Tensor, # (batch_size, num_features)
feat_static_real: torch.Tensor,
# (batch_size, num_features, history_length)
past_time_feat: torch.Tensor,
past_target: torch.Tensor, # (batch_size, history_length, 1)
@@ -186,6 +187,7 @@ class TransformerNetwork(nn.Module):
# (batch_size, num_features + prod(target_shape))
static_feat = torch.cat((
embedded_cat,
feat_static_real,
torch.log(scale)
if len(self.target_shape) == 0
else torch.log(scale.squeeze(1))),
@@ -214,6 +216,10 @@ class TransformerNetwork(nn.Module):
@staticmethod
def upper_triangular_mask(d):
return torch.triu(torch.ones((d,d)))
# mask = torch.zeros_like(torch.eye(d))
# for k in range(d - 1):
# mask = mask + torch.eye(d, d, k + 1)
# return mask
class TransformerTrainingNetwork(TransformerNetwork):
@@ -221,6 +227,7 @@ class TransformerTrainingNetwork(TransformerNetwork):
def forward(
self,
feat_static_cat: torch.Tensor,
feat_static_real: torch.Tensor,
past_time_feat: torch.Tensor,
past_target: torch.Tensor,
past_observed_values: torch.Tensor,
@@ -232,6 +239,7 @@ class TransformerTrainingNetwork(TransformerNetwork):
Parameters
----------
feat_static_cat : (batch_size, num_features)
feat_static_real: torch.Tensor, # (batch_size, num_features)
past_time_feat : (batch_size, history_length, num_features)
past_target : (batch_size, history_length, *target_shape)
past_observed_values : (batch_size, history_length, *target_shape, seq_len)
@@ -245,6 +253,7 @@ class TransformerTrainingNetwork(TransformerNetwork):
# create the inputs for the encoder
inputs, scale, _ = self.create_network_input(
feat_static_cat=feat_static_cat,
feat_static_real=feat_static_real,
past_time_feat=past_time_feat,
past_target=past_target,
past_observed_values=past_observed_values,
@@ -403,6 +412,8 @@ class TransformerPredictionNetwork(TransformerNetwork):
def forward(
self,
feat_static_cat: torch.Tensor,
feat_static_real: torch.Tensor,
feature_static_real: torch.Tensor,
past_time_feat: torch.Tensor,
past_target: torch.Tensor,
past_observed_values: torch.Tensor,
@@ -413,6 +424,7 @@ class TransformerPredictionNetwork(TransformerNetwork):
Parameters
----------
feat_static_cat : (batch_size, num_features)
feat_static_real : (batch_size, num_features)
past_time_feat : (batch_size, history_length, num_features)
past_target : (batch_size, history_length, *target_shape)
past_observed_values : (batch_size, history_length, *target_shape)
@@ -424,6 +436,7 @@ class TransformerPredictionNetwork(TransformerNetwork):
# create the inputs for the encoder
inputs, scale, static_feat = self.create_network_input(
feat_static_cat=feat_static_cat,
feature_static_real=feat_static_real,
past_time_feat=past_time_feat,
past_target=past_target,
past_observed_values=past_observed_values,