Getting started
This example trains a two-dimensional Real NVP model on a synthetic mixture of Gaussians. It includes data generation, optimization, density evaluation, and sampling. No downloaded dataset is needed.
Build and train a model
Run the following block as a Python script or in a notebook after completing Installation:
import torch
import antsnormflows as nf
torch.manual_seed(42)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def make_model():
base = nf.distributions.base.DiagGaussian(2)
flows = []
for _ in range(4):
conditioner = nf.nets.MLP([1, 64, 64, 2], init_zeros=True)
flows.append(nf.flows.AffineCouplingBlock(conditioner))
flows.append(nf.flows.Permute(2, mode="swap"))
return nf.NormalizingFlow(base, flows)
def draw_data(n):
centers = torch.tensor([[-2.0, 0.0], [2.0, 0.0]], device=device)
labels = torch.randint(2, (n,), device=device)
return centers[labels] + 0.5 * torch.randn(n, 2, device=device)
model = make_model().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
model.train()
for step in range(200):
x = draw_data(256)
optimizer.zero_grad(set_to_none=True)
loss = model.forward_kld(x)
if not torch.isfinite(loss):
raise RuntimeError("Non-finite training loss")
loss.backward()
optimizer.step()
if (step + 1) % 50 == 0:
print(f"Step {step + 1}: NLL = {loss.item():.3f}")
Each coupling layer leaves one coordinate unchanged and uses it to predict the scale and shift of the other coordinate. The conditioner therefore has one input and two outputs. Permutations exchange the coordinates so that both can be transformed across successive layers.
forward_kld(x) returns the mean negative log-likelihood of the batch.
Minimizing it fits the model to the observations; its value can be negative
because a continuous probability density can exceed one. This short run is
a demonstration, not a convergence guarantee.
Evaluate and sample
Continue in the same Python session:
model.eval()
with torch.no_grad():
validation = draw_data(1024)
log_prob = model.log_prob(validation)
samples, sample_log_prob = model.sample(512)
z = model.inverse(validation)
reconstructed = model(z)
print("Validation NLL:", -log_prob.mean().item())
print("Sample shape:", samples.shape) # (512, 2)
print("Log-density shape:", sample_log_prob.shape) # (512,)
print("Round-trip error:", (validation - reconstructed).abs().max().item())
log_prob gives one log density per observation, summed over its features.
sample returns both the generated observations and their model log
densities. model(z) maps latent coordinates to data, whereas
model.inverse(x) maps data to latent coordinates. Numerical round-trip
errors depend on the learned transformation and floating-point precision.
Save and reload
Continue with the same make_model function:
model.save("real_nvp.pt")
restored = make_model()
restored.load("real_nvp.pt", map_location="cpu")
restored.eval()
with torch.no_grad():
restored_samples, _ = restored.sample(16)
The checkpoint contains the model’s state dictionary, including learned parameters and registered buffers. Recreate the same architecture before loading it. This convenience method does not save the optimizer or training step; see Training and model conventions for a resumable checkpoint.