mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-09-09 11:32:17 +08:00
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:
committed by
GitHub Enterprise
parent
f701c90b37
commit
7a16872e26
+45
-75
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user