flow-pocket-lab / model.py
ARotting's picture
Publish Interactive temperature-controlled RealNVP sampler
a193678 verified
Raw
History Blame Contribute Delete
2.82 kB
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
@torch.inference_mode()
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())