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
This commit is contained in:
Kashif Rasul
2020-02-26 11:02:16 +01:00
committed by GitHub Enterprise
parent f701c90b37
commit 7a16872e26
+45 -75
View File
@@ -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