mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-09-10 12:37:53 +08:00
intial realnvp output distribution class
This commit is contained in:
@@ -19,6 +19,7 @@ from torch.distributions import (
|
||||
)
|
||||
|
||||
from .lambda_layer import LambdaLayer
|
||||
from .flows import RealNVP
|
||||
|
||||
|
||||
class ArgProj(nn.Module):
|
||||
@@ -88,7 +89,7 @@ class DistributionOutput(Output, ABC):
|
||||
|
||||
class BetaOutput(DistributionOutput):
|
||||
args_dim: Dict[str, int] = {"concentration1": 1, "concentration0": 1}
|
||||
distr_cls: Distribution = Beta
|
||||
distr_cls: type = Beta
|
||||
|
||||
@classmethod
|
||||
def domain_map(cls, concentration1, concentration0):
|
||||
@@ -236,3 +237,29 @@ class MultivariateNormalOutput(DistributionOutput):
|
||||
@property
|
||||
def event_shape(self) -> Tuple:
|
||||
return (self.dim,)
|
||||
|
||||
|
||||
|
||||
class RealNVPOutput(DistributionOutput):
|
||||
def __init__(self, input_size, cond_size, n_blocks=3, n_hidden=1, hidden_size=100):
|
||||
self.args_dim = {"cond": cond_label_size}
|
||||
self.dim = input_size
|
||||
self.flow = RealNVP(nblocks=nblocks,
|
||||
input_size=input_size,
|
||||
hidden_size=hidden_size,
|
||||
n_hidden=n_hidden,
|
||||
cond_label_size=cond_size
|
||||
)
|
||||
|
||||
def domain_map(self, cond):
|
||||
return cond
|
||||
|
||||
def distribution(self, distr_args, scale=None):
|
||||
cond, = distr_args
|
||||
self.flow.cond = cond
|
||||
|
||||
return self.flow
|
||||
|
||||
@property
|
||||
def event_shape(self) -> Tuple:
|
||||
return (self.dim,)
|
||||
+28
-9
@@ -42,7 +42,7 @@ class BatchNorm(nn.Module):
|
||||
if self.training:
|
||||
self.batch_mean = x.view(-1, x.shape[-1]).mean(0)
|
||||
# note MAF paper uses biased variance estimate; ie x.var(0, unbiased=False)
|
||||
self.batch_var = x.view(-1,x.shape[-1]).var(0)
|
||||
self.batch_var = x.view(-1, x.shape[-1]).var(0)
|
||||
|
||||
# update running mean
|
||||
self.running_mean.mul_(self.momentum).add_(
|
||||
@@ -142,12 +142,16 @@ class RealNVP(nn.Module):
|
||||
self.register_buffer('base_dist_mean', torch.zeros(input_size))
|
||||
self.register_buffer('base_dist_var', torch.ones(input_size))
|
||||
|
||||
self.__cond = 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)]
|
||||
hidden_size,
|
||||
n_hidden, mask,
|
||||
cond_label_size)]
|
||||
mask = 1 - mask
|
||||
modules += batch_norm * [BatchNorm(input_size)]
|
||||
|
||||
@@ -157,12 +161,27 @@ class RealNVP(nn.Module):
|
||||
def base_dist(self):
|
||||
return Normal(self.base_dist_mean, self.base_dist_var)
|
||||
|
||||
def forward(self, x, y=None):
|
||||
return self.net(x, y)
|
||||
@property
|
||||
def cond(self):
|
||||
return self.__cond
|
||||
|
||||
def inverse(self, u, y=None):
|
||||
return self.net.inverse(u, y)
|
||||
@cond.setter
|
||||
def cond(self, cond):
|
||||
self.__cond = cond
|
||||
|
||||
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).mean()
|
||||
def forward(self, x):
|
||||
return self.net(x, self.cond)
|
||||
|
||||
def inverse(self, u):
|
||||
return self.net.inverse(u, self.cond)
|
||||
|
||||
def log_prob(self, x):
|
||||
u, sum_log_abs_det_jacobians = self.forward(x, self.cond)
|
||||
return torch.sum(self.base_dist.log_prob(u) + sum_log_abs_det_jacobians, dim=-1)
|
||||
|
||||
def sample(sample_shape=None):
|
||||
|
||||
u = self.base_dist.sample(sample_shape)
|
||||
sample, _ = self.net.inverse(u, self.cond)
|
||||
|
||||
return sample
|
||||
Reference in New Issue
Block a user