From 68e91037ae048c11566045c74e5d4231f89bdc28 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Wed, 30 Oct 2019 12:46:45 +0100 Subject: [PATCH] added initial studentT output --- pts/model/deepar/deepar_estimator.py | 1 + pts/modules/__init__.py | 4 ++-- pts/modules/distribution_output.py | 14 +++++++++++++- 3 files changed, 16 insertions(+), 3 deletions(-) diff --git a/pts/model/deepar/deepar_estimator.py b/pts/model/deepar/deepar_estimator.py index 7aa1606..a3b602c 100644 --- a/pts/model/deepar/deepar_estimator.py +++ b/pts/model/deepar/deepar_estimator.py @@ -9,6 +9,7 @@ from pts.feature import ( time_features_from_frequency_str, ) from pts.model import PTSEstimator +from pts.modules import DistributionOutput, StudentTOutput class DeepAREstimator(PTSEstimator): diff --git a/pts/modules/__init__.py b/pts/modules/__init__.py index aac0cad..bc468ac 100644 --- a/pts/modules/__init__.py +++ b/pts/modules/__init__.py @@ -1,2 +1,2 @@ -from pts.modules.distribution_output import ArgProj -from pts.modules.lambda_layer import LambdaLayer +from .distribution_output import ArgProj, Output, DistributionOutput, StudentTOutput +from .lambda_layer import LambdaLayer diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index 39a4c0f..d38e7d4 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -4,7 +4,8 @@ from typing import Callable, Dict, Optional, Tuple import numpy as np import torch import torch.nn as nn -from torch.distributions import Distribution, TransformedDistribution, AffineTransform +import torch.nn.functional as F +from torch.distributions import * from .lambda_layer import LambdaLayer @@ -73,3 +74,14 @@ class DistributionOutput(Output): else: distr = self.distr_cls(*distr_args) return TransformedDistribution(distr, [AffineTransform(loc=0, scale=scale)]) + + +class StudentTOutput(DistributionOutput): + args_dim: Dict[str, int] = {"df": 1, "loc": 1, "scale": 1} + distr_cls: type = StudentT + + @classmethod + def domain_map(cls, df, loc, scale): + scale = nn.Softplus(scale) + df = 2.0 + nn.Softplus(df) + return df.squeeze(-1), loc.squeeze(-1), scale.squeeze(-1)