From ca2918b02e13c72a2b79e0521e25da8d88203f50 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Mon, 13 Jan 2020 16:59:07 +0100 Subject: [PATCH] fix comparison --- pts/modules/flows.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/pts/modules/flows.py b/pts/modules/flows.py index cea5922..ad5c442 100644 --- a/pts/modules/flows.py +++ b/pts/modules/flows.py @@ -180,13 +180,13 @@ class RealNVP(nn.Module): self.__scale = scale def forward(self, x): - if self.scale: + if self.scale is not None: x /= self.scale return self.net(x, self.cond) def inverse(self, u): x, log_abs_det_jacobian = self.net.inverse(u, self.cond) - if self.scale: + if self.scale is not None: x *= scale return x, log_abs_det_jacobian @@ -195,10 +195,12 @@ class RealNVP(nn.Module): return torch.sum(self.base_dist.log_prob(u) + sum_log_abs_det_jacobians, dim=-1) def sample(self, sample_shape=torch.Size()): - if self.cond: + if self.cond is not None: shape = self.cond.shape[:-1] + (self.input_size) - if sample_shape is not None: + else: shape = sample_shape + (self.input_size) + # if len(sample_shape) > 0: + # shape = sample_shape + (self.input_size) u = self.base_dist.sample(shape) sample, _ = self.inverse(u)