This commit is contained in:
Dr. Kashif Rasul
2020-01-04 12:23:08 +01:00
parent f050f80feb
commit 71fff3c411
3 changed files with 9 additions and 19 deletions
+1
View File
@@ -0,0 +1 @@
from .deepvar_estimator import DeepVAREstimator
+3 -3
View File
@@ -1,6 +1,5 @@
from typing import List, Optional
import numpy as np
import pandas as pd
import torch
@@ -27,6 +26,7 @@ from pts.transform import (
TargetDimIndicator,
)
from pts.feature import (
TimeFeature,
fourier_time_features_from_frequency_str,
get_fourier_lags_for_frequency,
)
@@ -85,7 +85,7 @@ class DeepVAREstimator(PTSEstimator):
self.cardinality = cardinality
self.embedding_dimension = embedding_dimension
self.conditioning_length = conditioning_length
self.use_marginal_t
self.use_marginal_transformation = use_marginal_transformation
self.lags_seq = (
lags_seq
@@ -204,7 +204,7 @@ class DeepVAREstimator(PTSEstimator):
transformation: Transformation,
trained_network: DeepVARTrainingNetwork,
device: torch.device,
) -> Predictor:
) -> PTSPredictor:
prediction_network = DeepVARPredictionNetwork(
input_size=self.input_size,
target_dim=self.target_dim,
+5 -16
View File
@@ -64,7 +64,7 @@ class DeepVARTrainingNetwork(nn.Module):
batch_first=True,
)
self.proj_dist_args = distr_output.get_args_proj()
self.proj_dist_args = distr_output.get_args_proj(num_cells)
self.embed_dim = 1
self.embed = nn.Embedding(
@@ -72,9 +72,9 @@ class DeepVARTrainingNetwork(nn.Module):
)
if scaling:
self.scaler = MeanScaler(keepdims=True)
self.scaler = MeanScaler(keepdim=True)
else:
self.scaler = NOPScaler(keepdims=True)
self.scaler = NOPScaler(keepdim=True)
@staticmethod
def get_lagged_subsequences(
@@ -118,7 +118,7 @@ class DeepVARTrainingNetwork(nn.Module):
begin_index = -lag_index - subsequences_length
end_index = -lag_index if lag_index > 0 else None
lagged_values.append(sequence[:, begin_index:end_index, ...].unsqueeze(1))
return torch.cat(*lagged_values, dim=1).permute(0, 2, 3, 1)
return torch.cat(lagged_values, dim=1).permute(0, 2, 3, 1)
def unroll(
self,
@@ -156,23 +156,12 @@ class DeepVARTrainingNetwork(nn.Module):
.expand(-1, unroll_length, -1)
.reshape((-1, unroll_length, self.target_dim * self.embed_dim))
)
# repeated_index_embeddings = (
# index_embeddings.expand_dims(axis=1)
# .repeat(axis=1, repeats=unroll_length)
# .reshape((-1, unroll_length, self.target_dim * self.embed_dim))
# )
# (batch_size, sub_seq_len, input_dim)
inputs = torch.cat((input_lags, repeated_index_embeddings, time_feat), dim=-1)
# unroll encoder
outputs, state = self.rnn(inputs, begin_state)
# inputs=inputs,
# length=unroll_length,
# layout="NTC",
# merge_outputs=True,
# begin_state=begin_state,
# )
# assert_shape(outputs, (-1, unroll_length, self.num_cells))
# for s in state:
@@ -412,7 +401,7 @@ class DeepVARTrainingNetwork(nn.Module):
# mask the loss at one time step if one or more observations is missing
# in the target dimensions (batch_size, subseq_length, 1)
loss_weights = observed_values.min(dim=-1, keepdims=True)
loss_weights = observed_values.min(dim=-1, keepdim=True)
# assert_shape(loss_weights, (-1, seq_len, 1))