From 98717c2e8f029bed52846e8a9a5ce25b7de6e442 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Mon, 6 Jan 2020 14:49:40 +0100 Subject: [PATCH] added multivariate gaussian output --- pts/modules/__init__.py | 1 + pts/modules/distribution_output.py | 27 ++++++++++++++++ test/modules/test_distribution_output.py | 39 ++++++++++++++++++++++++ 3 files changed, 67 insertions(+) diff --git a/pts/modules/__init__.py b/pts/modules/__init__.py index fd35a0b..59f493a 100644 --- a/pts/modules/__init__.py +++ b/pts/modules/__init__.py @@ -7,6 +7,7 @@ from .distribution_output import ( NegativeBinomialOutput, IndependentNormalOutput, LowRankMultivariateNormalOutput, + MultivariateNormalOutput, ) from .lambda_layer import LambdaLayer from .feature import FeatureEmbedder, FeatureAssembler diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index b26802a..453a569 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -198,3 +198,30 @@ class IndependentNormalOutput(DistributionOutput): return distr else: return TransformedDistribution(distr, [AffineTransform(loc=0, scale=scale)]) + + +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): + d = self.dim + device = scale.device + + shape = scale.shape[:-1] + (d, d) + scale = scale.reshape(shape) + + scale_diag = F.softplus(scale * torch.eye(d, device=device)) * torch.eye( + d, device=device + ) + + mask = torch.tril(torch.ones_like(scale), diagonal=-1) + scale_tril = (scale * mask) + scale_diag + + return loc, scale_tril + + @property + def event_shape(self) -> Tuple: + return (self.dim,) diff --git a/test/modules/test_distribution_output.py b/test/modules/test_distribution_output.py index 1576aab..b6e2658 100644 --- a/test/modules/test_distribution_output.py +++ b/test/modules/test_distribution_output.py @@ -13,6 +13,7 @@ from torch.distributions import ( Beta, NegativeBinomial, LowRankMultivariateNormal, + MultivariateNormal, ) from pts.modules import ( @@ -21,6 +22,7 @@ from pts.modules import ( BetaOutput, NegativeBinomialOutput, LowRankMultivariateNormalOutput, + MultivariateNormalOutput, ) NUM_SAMPLES = 2000 @@ -226,3 +228,40 @@ def test_lowrank_multivariate_normal() -> None: assert np.allclose( Sigma_hat, Sigma, atol=0.1, rtol=0.1 ), f"sigma did not match: sigma = {Sigma}, sigma_hat = {Sigma_hat}" + + +def test_multivariate_normal() -> None: + num_samples = 2000 + dim = 2 + + mu = np.arange(0, dim) / float(dim) + + L_diag = np.ones((dim,)) + L_low = 0.1 * np.ones((dim, dim)) * np.tri(dim, k=-1) + L = np.diag(L_diag) + L_low + Sigma = L.dot(L.transpose()) + + distr = MultivariateNormal(loc=torch.Tensor(mu), scale_tril=torch.Tensor(L)) + + samples = distr.sample((num_samples,)) + + mu_hat, L_hat = maximum_likelihood_estimate_sgd( + MultivariateNormalOutput(dim=dim), + samples, + init_biases=None, # todo we would need to rework biases a bit to use it in the multivariate case + learning_rate=0.01, + num_epochs=10, + ) + + distr = MultivariateNormal( + loc=torch.tensor(mu_hat), scale_tril=torch.tensor(L_hat) + ) + + Sigma_hat = distr.covariance_matrix.numpy() + + assert np.allclose( + mu_hat, mu, atol=0.1, rtol=0.1 + ), f"mu did not match: mu = {mu}, mu_hat = {mu_hat}" + assert np.allclose( + Sigma_hat, Sigma, atol=0.1, rtol=0.1 + ), f"Sigma did not match: sigma = {Sigma}, sigma_hat = {Sigma_hat}" \ No newline at end of file