From 7a16872e26474131ddd32e0c888732fe1ab23589 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Wed, 26 Feb 2020 11:02:16 +0100 Subject: [PATCH] fix the final log_abs_det_jacobian when scaling (#3) * fix the final log_abs_det_jacobian * fix argument name * subclass from Flow class * move buffer for base dist to Flow class as well * formatting --- pts/modules/flows.py | 120 ++++++++++++++++--------------------------- 1 file changed, 45 insertions(+), 75 deletions(-) diff --git a/pts/modules/flows.py b/pts/modules/flows.py index 815f266..582e8df 100644 --- a/pts/modules/flows.py +++ b/pts/modules/flows.py @@ -162,8 +162,10 @@ class LinearMaskedCoupling(nn.Module): # run through model s = self.s_net(mx if y is None else torch.cat([y, mx], dim=-1)) - t = self.t_net(mx if y is None else torch.cat([y, mx], dim=-1)) * (1 - self.mask) - + t = self.t_net(mx if y is None else torch.cat([y, mx], dim=-1)) * ( + 1 - self.mask + ) + # cf RealNVP eq 8 where u corresponds to x (here we're modeling u) log_s = torch.tanh(s) * (1 - self.mask) u = x * torch.exp(log_s) + t @@ -183,7 +185,9 @@ class LinearMaskedCoupling(nn.Module): # run through model s = self.s_net(mu if y is None else torch.cat([y, mu], dim=-1)) - t = self.t_net(mu if y is None else torch.cat([y, mu], dim=-1)) * (1 - self.mask) + t = self.t_net(mu if y is None else torch.cat([y, mu], dim=-1)) * ( + 1 - self.mask + ) log_s = torch.tanh(s) * (1 - self.mask) x = (u - t) * torch.exp(-log_s) @@ -298,38 +302,15 @@ class MADE(nn.Module): return torch.sum(self.base_dist.log_prob(u) + log_abs_det_jacobian, dim=-1) -class RealNVP(nn.Module): - def __init__( - self, - n_blocks, - input_size, - hidden_size, - n_hidden, - cond_label_size=None, - batch_norm=True, - ): - super().__init__() +class Flow(nn.Module): + def __init__(self, input_size): + self.__scale = None + self.net = None # base distribution for calculation of log prob under the model self.register_buffer("base_dist_mean", torch.zeros(input_size)) self.register_buffer("base_dist_var", torch.ones(input_size)) - self.__scale = None - - # construct model - modules = [] - mask = torch.arange(input_size).float() % 2 - for i in range(n_blocks): - modules += [ - LinearMaskedCoupling( - input_size, hidden_size, n_hidden, mask, cond_label_size - ) - ] - mask = 1 - mask - modules += batch_norm * [BatchNorm(input_size)] - - self.net = FlowSequential(*modules) - @property def base_dist(self): return Normal(self.base_dist_mean, self.base_dist_var) @@ -345,12 +326,16 @@ class RealNVP(nn.Module): def forward(self, x, cond): if self.scale is not None: x /= self.scale - return self.net(x, cond) + u, log_abs_det_jacobian = self.net(x, cond) + if self.scale is not None: + log_abs_det_jacobian -= torch.log(torch.abs(self.scale)) + return u, log_abs_det_jacobian def inverse(self, u, cond): x, log_abs_det_jacobian = self.net.inverse(u, cond) if self.scale is not None: x *= self.scale + log_abs_det_jacobian += torch.log(torch.abs(self.scale)) return x, log_abs_det_jacobian def log_prob(self, x, cond): @@ -368,8 +353,34 @@ class RealNVP(nn.Module): return sample +class RealNVP(Flow): + def __init__( + self, + n_blocks, + input_size, + hidden_size, + n_hidden, + cond_label_size=None, + batch_norm=True, + ): + super().__init__(input_size) -class MAF(nn.Module): + # construct model + modules = [] + mask = torch.arange(input_size).float() % 2 + for i in range(n_blocks): + modules += [ + LinearMaskedCoupling( + input_size, hidden_size, n_hidden, mask, cond_label_size + ) + ] + mask = 1 - mask + modules += batch_norm * [BatchNorm(input_size)] + + self.net = FlowSequential(*modules) + + +class MAF(Flow): def __init__( self, n_blocks, @@ -381,12 +392,8 @@ class MAF(nn.Module): input_order="sequential", batch_norm=True, ): - super().__init__() - # base distribution for calculation of log prob under the model - self.register_buffer("base_dist_mean", torch.zeros(input_size)) - self.register_buffer("base_dist_var", torch.ones(input_size)) - self.__scale = None - + super().__init__(input_size) + # construct model modules = [] self.input_degrees = None @@ -406,40 +413,3 @@ class MAF(nn.Module): modules += batch_norm * [BatchNorm(input_size)] self.net = FlowSequential(*modules) - - @property - def base_dist(self): - return Normal(self.base_dist_mean, self.base_dist_var) - - @property - def scale(self): - return self.__scale - - @scale.setter - def scale(self, scale): - self.__scale = scale - - def forward(self, x, y=None): - if self.scale is not None: - x /= self.scale - return self.net(x, y) - - def inverse(self, u, y=None): - x, log_abs_det_jacobian = self.net.inverse(u, y) - if self.scale is not None: - x *= self.scale - return x, log_abs_det_jacobian - - def log_prob(self, x, y=None): - u, sum_log_abs_det_jacobians = self.forward(x, y) - return torch.sum(self.base_dist.log_prob(u) + sum_log_abs_det_jacobians, dim=-1) - - def sample(self, sample_shape=torch.Size(), cond=None): - if cond is not None: - shape = cond.shape[:-1] - else: - shape = sample_shape - - u = self.base_dist.sample(shape) - sample, _ = self.inverse(u, cond) - return sample