Spaces:
Running
Running
| from __future__ import annotations | |
| import math | |
| import torch | |
| from torch import nn | |
| class CouplingLayer(nn.Module): | |
| def __init__(self, mask: tuple[float, float]) -> None: | |
| super().__init__() | |
| self.register_buffer("mask", torch.tensor(mask)) | |
| self.network = nn.Sequential( | |
| nn.Linear(2, 48), | |
| nn.SiLU(), | |
| nn.Linear(48, 48), | |
| nn.SiLU(), | |
| nn.Linear(48, 4), | |
| ) | |
| def parameters_for(self, masked: torch.Tensor) -> tuple[torch.Tensor, ...]: | |
| scale, translation = self.network(masked).chunk(2, dim=1) | |
| scale = 1.4 * torch.tanh(scale) * (1 - self.mask) | |
| translation = translation * (1 - self.mask) | |
| return scale, translation | |
| def forward(self, values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| masked = values * self.mask | |
| scale, translation = self.parameters_for(masked) | |
| transformed = masked + (1 - self.mask) * ( | |
| values * torch.exp(scale) + translation | |
| ) | |
| return transformed, scale.sum(1) | |
| def inverse(self, values: torch.Tensor) -> torch.Tensor: | |
| masked = values * self.mask | |
| scale, translation = self.parameters_for(masked) | |
| return masked + (1 - self.mask) * ( | |
| (values - translation) * torch.exp(-scale) | |
| ) | |
| class RealNVP(nn.Module): | |
| def __init__(self, layers: int = 8) -> None: | |
| super().__init__() | |
| self.layers = nn.ModuleList( | |
| [ | |
| CouplingLayer((1.0, 0.0) if index % 2 == 0 else (0.0, 1.0)) | |
| for index in range(layers) | |
| ] | |
| ) | |
| def forward(self, values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| log_determinant = values.new_zeros(len(values)) | |
| latent = values | |
| for layer in self.layers: | |
| latent, change = layer(latent) | |
| log_determinant += change | |
| return latent, log_determinant | |
| def inverse(self, latent: torch.Tensor) -> torch.Tensor: | |
| values = latent | |
| for layer in reversed(self.layers): | |
| values = layer.inverse(values) | |
| return values | |
| def log_probability(self, values: torch.Tensor) -> torch.Tensor: | |
| latent, log_determinant = self(values) | |
| base = -0.5 * (latent**2).sum(1) - math.log(2 * math.pi) | |
| return base + log_determinant | |
| def sample( | |
| self, | |
| samples: int, | |
| *, | |
| seed: int, | |
| temperature: float = 1.0, | |
| ) -> torch.Tensor: | |
| generator = torch.Generator().manual_seed(seed) | |
| latent = temperature * torch.randn(samples, 2, generator=generator) | |
| return self.inverse(latent) | |
| def parameter_count(module: nn.Module) -> int: | |
| return sum(parameter.numel() for parameter in module.parameters()) | |