From e263bd5f5c4a6df0b3420958fa826351f1fbddc7 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Thu, 9 Jan 2020 13:46:56 +0100 Subject: [PATCH] fix past_observed_values --- pts/model/deepvar/deepvar_network.py | 2 +- pts/modules/distribution_output.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/pts/model/deepvar/deepvar_network.py b/pts/model/deepvar/deepvar_network.py index 37cc7f3..c3a187d 100644 --- a/pts/model/deepvar/deepvar_network.py +++ b/pts/model/deepvar/deepvar_network.py @@ -234,7 +234,7 @@ class DeepVARTrainingNetwork(nn.Module): """ past_observed_values = torch.min( - past_observed_values, past_is_pad.unsqueeze(-1) + past_observed_values, 1 - past_is_pad.unsqueeze(-1) ) if future_time_feat is None or future_target_cdf is None: diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index 515a1a6..b4d2dee 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -203,7 +203,6 @@ class IndependentNormalOutput(DistributionOutput): class MultivariateNormalOutput(DistributionOutput): def __init__(self, dim: int) -> None: self.args_dim = {"loc": dim, "scale_tril": dim * dim} - self.distr_cls = MultivariateNormal self.dim = dim def domain_map(self, loc, scale):