mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-08-20 12:40:22 +08:00
added static real features
This commit is contained in:
@@ -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,
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Reference in New Issue
Block a user