From e81d2448329c4d435453a249dce5ff14f3190606 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Thu, 21 Nov 2019 23:13:38 +0100 Subject: [PATCH] get args takes in in_feature argument --- pts/model/deepar/deepar_network.py | 2 +- pts/modules/distribution_output.py | 10 +++++++--- test/modules/test_distribution_output.py | 3 +-- 3 files changed, 9 insertions(+), 6 deletions(-) diff --git a/pts/model/deepar/deepar_network.py b/pts/model/deepar/deepar_network.py index e9b285c..4d4eae8 100644 --- a/pts/model/deepar/deepar_network.py +++ b/pts/model/deepar/deepar_network.py @@ -53,7 +53,7 @@ class DeepARNetwork(nn.Module): # TODO # self.target_shape = distr_output.event_shape - self.proj_distr_args = distr_output.get_args_proj() + self.proj_distr_args = distr_output.get_args_proj(num_cells) self.embedder = FeatureEmbedder(cardinalities=cardinality, embedding_dims=embedding_dimension) diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index 12f6799..3859797 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -5,7 +5,7 @@ import numpy as np import torch import torch.nn as nn import torch.nn.functional as F -from torch.distributions import * +from torch.distributions import Distribution, StudentT, TransformedDistribution, AffineTransform from .lambda_layer import LambdaLayer @@ -47,9 +47,9 @@ class Output(ABC): def dtype(self, dtype: np.dtype): self._dtype = dtype - def get_args_proj(self, prefix: Optional[str] = None) -> ArgProj: + def get_args_proj(self, in_features: int, prefix: Optional[str] = None) -> ArgProj: return ArgProj( - in_features=self.in_features, + in_features=in_features, args_dim=self.args_dim, domain_map=LambdaLayer(self.domain_map), prefix=prefix, @@ -85,3 +85,7 @@ class StudentTOutput(DistributionOutput): scale = F.softplus(scale) df = 2.0 + F.softplus(df) return df.squeeze(-1), loc.squeeze(-1), scale.squeeze(-1) + + @property + def event_shape(self) -> Tuple: + return () \ No newline at end of file diff --git a/test/modules/test_distribution_output.py b/test/modules/test_distribution_output.py index 5beba8c..de1ad8c 100644 --- a/test/modules/test_distribution_output.py +++ b/test/modules/test_distribution_output.py @@ -27,8 +27,7 @@ def maximum_likelihood_estimate_sgd(distr_output: DistributionOutput, init_biases: List[np.ndarray] = None, num_epochs: int = 5, learning_rate: float = 1e-2): - distr_output.in_features = 1 - arg_proj = distr_output.get_args_proj() + arg_proj = distr_output.get_args_proj(in_features=1) if init_biases is not None: for param, bias in zip(arg_proj.proj, init_biases):