diff --git a/pts/modules/distribution_output.py b/pts/modules/distribution_output.py index b4d2dee..646d2b6 100644 --- a/pts/modules/distribution_output.py +++ b/pts/modules/distribution_output.py @@ -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,) \ No newline at end of file diff --git a/pts/modules/flows.py b/pts/modules/flows.py index 66f69ae..d2e893a 100644 --- a/pts/modules/flows.py +++ b/pts/modules/flows.py @@ -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 \ No newline at end of file