From 55644c46714cfc097fe4bbc6c831452917ba1749 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Wed, 25 Mar 2020 13:34:08 +0100 Subject: [PATCH] added NormalOutput --- pts/modules/__init__.py | 1 + pts/modules/distribution_output.py | 15 +++++++++++++++ 2 files changed, 16 insertions(+) diff --git a/pts/modules/__init__.py b/pts/modules/__init__.py index 10c46c9..73090f8 100644 --- a/pts/modules/__init__.py +++ b/pts/modules/__init__.py @@ -2,6 +2,7 @@ from .distribution_output import ( ArgProj, Output, DistributionOutput, + NormalOutput, StudentTOutput, BetaOutput, NegativeBinomialOutput, diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index dff3eb9..56e462c 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -8,6 +8,7 @@ import torch.nn.functional as F from pts.core.component import validated from torch.distributions import ( Distribution, + Normal, Beta, NegativeBinomial, StudentT, @@ -93,6 +94,20 @@ class DistributionOutput(Output, ABC): return TransformedDistribution(distr, [AffineTransform(loc=0, scale=scale)]) +class NormalOutput(DistributionOutput): + args_dim: Dict[str, int] = {"loc": 1, "scale": 1} + distr_cls: type = Normal + + @classmethod + def domain_map(self, loc, scale): + scale = F.softplus(scale) + return loc.squeeze(-1), scale.squeeze(-1) + + @property + def event_shape(self) -> Tuple: + return () + + class BetaOutput(DistributionOutput): args_dim: Dict[str, int] = {"concentration1": 1, "concentration0": 1} distr_cls: type = Beta