From af0b129b2120ee6e12670fbafc4dcab7c30fd750 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Fri, 15 May 2020 13:12:44 +0200 Subject: [PATCH] fixed typos --- pts/modules/distribution_output.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index 66f4345..ae2166d 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -247,8 +247,7 @@ class LowRankMultivariateNormalOutput(DistributionOutput): self.sigma_minimum = sigma_minimum self.args_dim = {"loc": dim, "cov_factor": dim * rank, "cov_diag": dim} - @classmethod - def domain_map(cls, loc, cov_factor, cov_diag): + def domain_map(self, loc, cov_factor, cov_diag): diag_bias = ( self.inv_softplus(self.sigma_init ** 2) if self.sigma_init > 0.0 else 0.0 ) @@ -299,9 +298,8 @@ class MultivariateNormalOutput(DistributionOutput): self.args_dim = {"loc": dim, "scale_tril": dim * dim} self.dim = dim - @classmethod - def domain_map(cls, loc, scale): - d = len(loc) + def domain_map(self, loc, scale): + d = self.dim device = scale.device shape = scale.shape[:-1] + (d, d)