Synthetic models for posterior distributions¶
Marco Raveri (marco.raveri@unige.it), Cyrille Doux (doux@lpsc.in2p3.fr), Shivam Pandey (shivampcosmo@gmail.com)
In this notebook we show how to build normalizing flow syntetic models for posterior distributions, as in Raveri, Doux and Pandey (2024), arXiv:2409.09101.
Table of contents¶
Notebook setup:¶
# Show plots inline, and load main getdist plot module and samples class
%matplotlib inline
%config InlineBackend.figure_format = 'retina'
%load_ext autoreload
%autoreload 2
# import libraries:
import sys, os
os.environ['TF_USE_LEGACY_KERAS'] = '1' # needed for tensorflow KERAS compatibility
os.environ['DISPLAY'] = 'inline' # hack to get getdist working
sys.path.insert(0,os.path.realpath(os.path.join(os.getcwd(),'../..')))
from getdist import plots, MCSamples
from getdist.gaussian_mixtures import GaussianND
import getdist
getdist.chains.print_load_details = False
import scipy
import matplotlib.pyplot as plt
import numpy as np
import seaborn as sns
# tensorflow imports:
import tensorflow as tf
import tensorflow_probability as tfp
# import the tensiometer tools that we need:
import tensiometer
from tensiometer.utilities import stats_utilities as utilities
from tensiometer.synthetic_probability import synthetic_probability as sp
# getdist settings to ensure consistency of plots:
getdist_settings = {'ignore_rows': 0.0,
'smooth_scale_2D': 0.3,
'smooth_scale_1D': 0.3,
}
2026-01-05 23:41:03.608237: I tensorflow/core/platform/cpu_feature_guard.cc:210] This TensorFlow binary is optimized to use available CPU instructions in performance-critical operations. To enable the following instructions: AVX2 FMA, in other operations, rebuild TensorFlow with the appropriate compiler flags.
We start by building a random Gaussian mixture that we are going to use for tests:
# define the parameters of the problem:
dim = 6
num_gaussians = 3
num_samples = 10000
# we seed the random number generator to get reproducible results:
seed = 100
np.random.seed(seed)
# we define the range for the means and covariances:
mean_range = (-0.5, 0.5)
cov_scale = 0.4**2
# means and covs:
means = np.random.uniform(mean_range[0], mean_range[1], num_gaussians*dim).reshape(num_gaussians, dim)
weights = np.random.rand(num_gaussians)
weights = weights / np.sum(weights)
covs = [cov_scale*utilities.vector_to_PDM(np.random.rand(int(dim*(dim+1)/2))) for _ in range(num_gaussians)]
# cast to required precision:
means = means.astype(np.float32)
weights = weights.astype(np.float32)
covs = [cov.astype(np.float32) for cov in covs]
# initialize distribution:
distribution = tfp.distributions.Mixture(
cat=tfp.distributions.Categorical(probs=weights),
components=[
tfp.distributions.MultivariateNormalTriL(loc=_m, scale_tril=tf.linalg.cholesky(_c))
for _m, _c in zip(means, covs)
], name='Mixture')
# sample the distribution:
samples = distribution.sample(num_samples).numpy()
# calculate log posteriors:
logP = distribution.log_prob(samples).numpy()
# create MCSamples from the samples:
chain = MCSamples(samples=samples,
settings=getdist_settings,
loglikes=-logP,
name_tag='Mixture',
)
# we make a sanity check plot:
g = plots.get_subplot_plotter()
g.triangle_plot(chain, filled=True)
plt.show();
Base example:¶
We train a normalizing flow on samples of a given distribution.
We initialize and train the normalizing flow on samples of the distribution we have just defined:
kwargs = {
'feedback': 2,
'plot_every': 1000,
'pop_size': 1,
#'cache_dir': 'test', # set this to a directory to cache the results
#'root_name': 'test', # sets the name of the flow for the cache files
}
flow = sp.flow_from_chain(chain, # parameter difference chain
**kwargs)
* Initializing samples
- flow name: Mixture_flow
- precision: <dtype: 'float32'>
- flow parameters and ranges:
param1 : [-1.36772, 1.16794]
param2 : [-1.43056, 1.23563]
param3 : [-1.25734, 0.778261]
param4 : [-0.8318, 1.23481]
param5 : [-1.72142, 1.3057]
param6 : [-1.55835, 0.937404]
- periodic parameters: []
- time taken: 0.0027 seconds
* Initializing fixed bijector
- using prior bijector: ranges
- rescaling samples
- time taken: 0.7350 seconds
* Initializing trainable bijector
Building Autoregressive Flow
- # parameters : 6
- periodic parameters : None
- # transformations : 8
- hidden_units : [16, 16]
- transformation_type : affine
- autoregressive_type : masked
- permutations : True
- scale_roto_shift : False
- activation : <function asinh at 0x18f497060>
- time taken: 4.5673 seconds
* Initializing training dataset
- 9000/1000 training/test samples and uniform weights
- time taken: 4.5660 seconds
* Initializing transformed distribution
- time taken: 0.0351 seconds
* Initializing loss function
- using standard loss function
- time taken: 0.0001 seconds
* Initializing training model
- Compiling model
- time taken: 0.0924 seconds
- trainable parameters : 4704
- maximum learning rate: 0.001
- minimum learning rate: 1e-06
- time taken: 3.6294 seconds
* Training
- Compiling model
- time taken: 0.0686 seconds
Epoch 1/100
20/20 - 21s - loss: 8.4837 - val_loss: 8.5265 - lr: 0.0010 - 21s/epoch - 1s/step
Epoch 2/100
20/20 - 0s - loss: 8.4392 - val_loss: 8.4783 - lr: 0.0010 - 277ms/epoch - 14ms/step
Epoch 3/100
20/20 - 0s - loss: 8.3918 - val_loss: 8.4481 - lr: 0.0010 - 282ms/epoch - 14ms/step
Epoch 4/100
20/20 - 0s - loss: 8.3508 - val_loss: 8.4243 - lr: 0.0010 - 278ms/epoch - 14ms/step
Epoch 5/100
20/20 - 0s - loss: 8.3227 - val_loss: 8.4112 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 6/100
20/20 - 0s - loss: 8.3083 - val_loss: 8.4087 - lr: 0.0010 - 271ms/epoch - 14ms/step
Epoch 7/100
20/20 - 0s - loss: 8.3005 - val_loss: 8.4016 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 8/100
20/20 - 0s - loss: 8.2942 - val_loss: 8.3987 - lr: 0.0010 - 273ms/epoch - 14ms/step
Epoch 9/100
20/20 - 0s - loss: 8.2880 - val_loss: 8.3928 - lr: 0.0010 - 271ms/epoch - 14ms/step
Epoch 10/100
20/20 - 0s - loss: 8.2814 - val_loss: 8.3857 - lr: 0.0010 - 273ms/epoch - 14ms/step
Epoch 11/100
20/20 - 0s - loss: 8.2715 - val_loss: 8.3805 - lr: 0.0010 - 272ms/epoch - 14ms/step
Epoch 12/100
20/20 - 0s - loss: 8.2589 - val_loss: 8.3643 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 13/100
20/20 - 0s - loss: 8.2413 - val_loss: 8.3440 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 14/100
20/20 - 0s - loss: 8.2164 - val_loss: 8.3221 - lr: 0.0010 - 274ms/epoch - 14ms/step
Epoch 15/100
20/20 - 0s - loss: 8.1851 - val_loss: 8.2868 - lr: 0.0010 - 280ms/epoch - 14ms/step
Epoch 16/100
20/20 - 0s - loss: 8.1495 - val_loss: 8.2479 - lr: 0.0010 - 307ms/epoch - 15ms/step
Epoch 17/100
20/20 - 0s - loss: 8.1144 - val_loss: 8.2099 - lr: 0.0010 - 281ms/epoch - 14ms/step
Epoch 18/100
20/20 - 0s - loss: 8.0873 - val_loss: 8.1785 - lr: 0.0010 - 288ms/epoch - 14ms/step
Epoch 19/100
20/20 - 0s - loss: 8.0594 - val_loss: 8.1523 - lr: 0.0010 - 279ms/epoch - 14ms/step
Epoch 20/100
20/20 - 0s - loss: 8.0322 - val_loss: 8.1170 - lr: 0.0010 - 284ms/epoch - 14ms/step
Epoch 21/100
20/20 - 0s - loss: 8.0055 - val_loss: 8.0967 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 22/100
20/20 - 0s - loss: 7.9782 - val_loss: 8.0727 - lr: 0.0010 - 284ms/epoch - 14ms/step
Epoch 23/100
20/20 - 0s - loss: 7.9495 - val_loss: 8.0355 - lr: 0.0010 - 281ms/epoch - 14ms/step
Epoch 24/100
20/20 - 0s - loss: 7.9212 - val_loss: 8.0099 - lr: 0.0010 - 293ms/epoch - 15ms/step
Epoch 25/100
20/20 - 0s - loss: 7.8907 - val_loss: 7.9787 - lr: 0.0010 - 293ms/epoch - 15ms/step
Epoch 26/100
20/20 - 0s - loss: 7.8602 - val_loss: 7.9549 - lr: 0.0010 - 298ms/epoch - 15ms/step
Epoch 27/100
20/20 - 0s - loss: 7.8323 - val_loss: 7.9183 - lr: 0.0010 - 286ms/epoch - 14ms/step
Epoch 28/100
20/20 - 0s - loss: 7.8062 - val_loss: 7.9032 - lr: 0.0010 - 293ms/epoch - 15ms/step
Epoch 29/100
20/20 - 0s - loss: 7.7786 - val_loss: 7.8651 - lr: 0.0010 - 289ms/epoch - 14ms/step
Epoch 30/100
20/20 - 0s - loss: 7.7578 - val_loss: 7.8517 - lr: 0.0010 - 301ms/epoch - 15ms/step
Epoch 31/100
20/20 - 1s - loss: 7.7337 - val_loss: 7.8204 - lr: 0.0010 - 558ms/epoch - 28ms/step
Epoch 32/100
20/20 - 0s - loss: 7.7097 - val_loss: 7.7945 - lr: 0.0010 - 297ms/epoch - 15ms/step
Epoch 33/100
20/20 - 0s - loss: 7.6867 - val_loss: 7.7822 - lr: 0.0010 - 471ms/epoch - 24ms/step
Epoch 34/100
20/20 - 0s - loss: 7.6665 - val_loss: 7.7561 - lr: 0.0010 - 359ms/epoch - 18ms/step
Epoch 35/100
20/20 - 0s - loss: 7.6507 - val_loss: 7.7342 - lr: 0.0010 - 334ms/epoch - 17ms/step
Epoch 36/100
20/20 - 0s - loss: 7.6324 - val_loss: 7.7147 - lr: 0.0010 - 316ms/epoch - 16ms/step
Epoch 37/100
20/20 - 0s - loss: 7.6228 - val_loss: 7.7019 - lr: 0.0010 - 322ms/epoch - 16ms/step
Epoch 38/100
20/20 - 0s - loss: 7.6054 - val_loss: 7.6892 - lr: 0.0010 - 298ms/epoch - 15ms/step
Epoch 39/100
20/20 - 0s - loss: 7.5909 - val_loss: 7.6767 - lr: 0.0010 - 289ms/epoch - 14ms/step
Epoch 40/100
20/20 - 0s - loss: 7.5840 - val_loss: 7.6746 - lr: 0.0010 - 297ms/epoch - 15ms/step
Epoch 41/100
20/20 - 0s - loss: 7.5710 - val_loss: 7.6602 - lr: 0.0010 - 286ms/epoch - 14ms/step
Epoch 42/100
20/20 - 0s - loss: 7.5648 - val_loss: 7.6545 - lr: 0.0010 - 287ms/epoch - 14ms/step
Epoch 43/100
20/20 - 0s - loss: 7.5525 - val_loss: 7.6542 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 44/100
20/20 - 0s - loss: 7.5483 - val_loss: 7.6507 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 45/100
20/20 - 0s - loss: 7.5435 - val_loss: 7.6383 - lr: 0.0010 - 289ms/epoch - 14ms/step
Epoch 46/100
20/20 - 0s - loss: 7.5327 - val_loss: 7.6404 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 47/100
20/20 - 0s - loss: 7.5294 - val_loss: 7.6223 - lr: 0.0010 - 284ms/epoch - 14ms/step
Epoch 48/100
20/20 - 0s - loss: 7.5206 - val_loss: 7.6354 - lr: 0.0010 - 293ms/epoch - 15ms/step
Epoch 49/100
20/20 - 0s - loss: 7.5159 - val_loss: 7.6275 - lr: 0.0010 - 274ms/epoch - 14ms/step
Epoch 50/100
20/20 - 0s - loss: 7.5081 - val_loss: 7.6185 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 51/100
20/20 - 0s - loss: 7.5036 - val_loss: 7.6048 - lr: 0.0010 - 279ms/epoch - 14ms/step
Epoch 52/100
20/20 - 0s - loss: 7.4972 - val_loss: 7.6304 - lr: 0.0010 - 281ms/epoch - 14ms/step
Epoch 53/100
20/20 - 0s - loss: 7.4975 - val_loss: 7.6141 - lr: 0.0010 - 302ms/epoch - 15ms/step
Epoch 54/100
20/20 - 0s - loss: 7.4929 - val_loss: 7.6018 - lr: 0.0010 - 265ms/epoch - 13ms/step
Epoch 55/100
20/20 - 0s - loss: 7.4910 - val_loss: 7.5980 - lr: 0.0010 - 269ms/epoch - 13ms/step
Epoch 56/100
20/20 - 0s - loss: 7.4828 - val_loss: 7.5917 - lr: 0.0010 - 286ms/epoch - 14ms/step
Epoch 57/100
20/20 - 0s - loss: 7.4797 - val_loss: 7.6028 - lr: 0.0010 - 270ms/epoch - 13ms/step
Epoch 58/100
20/20 - 0s - loss: 7.4799 - val_loss: 7.5999 - lr: 0.0010 - 280ms/epoch - 14ms/step
Epoch 59/100
20/20 - 0s - loss: 7.4816 - val_loss: 7.5841 - lr: 0.0010 - 269ms/epoch - 13ms/step
Epoch 60/100
20/20 - 0s - loss: 7.4726 - val_loss: 7.5907 - lr: 0.0010 - 269ms/epoch - 13ms/step
Epoch 61/100
20/20 - 0s - loss: 7.4659 - val_loss: 7.5885 - lr: 0.0010 - 269ms/epoch - 13ms/step
Epoch 62/100
20/20 - 0s - loss: 7.4640 - val_loss: 7.5798 - lr: 0.0010 - 262ms/epoch - 13ms/step
Epoch 63/100
20/20 - 0s - loss: 7.4610 - val_loss: 7.5820 - lr: 0.0010 - 270ms/epoch - 14ms/step
Epoch 64/100
20/20 - 0s - loss: 7.4591 - val_loss: 7.5645 - lr: 0.0010 - 263ms/epoch - 13ms/step
Epoch 65/100
20/20 - 0s - loss: 7.4555 - val_loss: 7.5746 - lr: 0.0010 - 268ms/epoch - 13ms/step
Epoch 66/100
20/20 - 0s - loss: 7.4508 - val_loss: 7.5886 - lr: 0.0010 - 274ms/epoch - 14ms/step
Epoch 67/100
20/20 - 0s - loss: 7.4498 - val_loss: 7.5589 - lr: 0.0010 - 273ms/epoch - 14ms/step
Epoch 68/100
20/20 - 0s - loss: 7.4441 - val_loss: 7.5658 - lr: 0.0010 - 270ms/epoch - 13ms/step
Epoch 69/100
20/20 - 0s - loss: 7.4421 - val_loss: 7.5657 - lr: 0.0010 - 272ms/epoch - 14ms/step
Epoch 70/100
20/20 - 0s - loss: 7.4391 - val_loss: 7.5670 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 71/100
20/20 - 0s - loss: 7.4420 - val_loss: 7.5536 - lr: 0.0010 - 271ms/epoch - 14ms/step
Epoch 72/100
20/20 - 0s - loss: 7.4395 - val_loss: 7.5465 - lr: 0.0010 - 304ms/epoch - 15ms/step
Epoch 73/100
20/20 - 0s - loss: 7.4330 - val_loss: 7.5533 - lr: 0.0010 - 276ms/epoch - 14ms/step
Epoch 74/100
20/20 - 0s - loss: 7.4343 - val_loss: 7.5467 - lr: 0.0010 - 274ms/epoch - 14ms/step
Epoch 75/100
20/20 - 0s - loss: 7.4288 - val_loss: 7.5488 - lr: 0.0010 - 296ms/epoch - 15ms/step
Epoch 76/100
20/20 - 0s - loss: 7.4251 - val_loss: 7.5448 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 77/100
20/20 - 0s - loss: 7.4238 - val_loss: 7.5408 - lr: 0.0010 - 277ms/epoch - 14ms/step
Epoch 78/100
20/20 - 0s - loss: 7.4217 - val_loss: 7.5426 - lr: 0.0010 - 278ms/epoch - 14ms/step
Epoch 79/100
20/20 - 0s - loss: 7.4150 - val_loss: 7.5441 - lr: 0.0010 - 287ms/epoch - 14ms/step
Epoch 80/100
20/20 - 0s - loss: 7.4132 - val_loss: 7.5453 - lr: 0.0010 - 288ms/epoch - 14ms/step
Epoch 81/100
20/20 - 0s - loss: 7.4163 - val_loss: 7.5429 - lr: 0.0010 - 285ms/epoch - 14ms/step
Epoch 82/100
20/20 - 0s - loss: 7.4105 - val_loss: 7.5306 - lr: 0.0010 - 287ms/epoch - 14ms/step
Epoch 83/100
20/20 - 0s - loss: 7.4120 - val_loss: 7.5469 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 84/100
20/20 - 0s - loss: 7.4108 - val_loss: 7.5269 - lr: 0.0010 - 312ms/epoch - 16ms/step
Epoch 85/100
20/20 - 0s - loss: 7.4035 - val_loss: 7.5284 - lr: 0.0010 - 300ms/epoch - 15ms/step
Epoch 86/100
20/20 - 0s - loss: 7.4017 - val_loss: 7.5271 - lr: 0.0010 - 278ms/epoch - 14ms/step
Epoch 87/100
20/20 - 0s - loss: 7.3991 - val_loss: 7.5337 - lr: 0.0010 - 290ms/epoch - 15ms/step
Epoch 88/100
20/20 - 0s - loss: 7.3986 - val_loss: 7.5215 - lr: 0.0010 - 289ms/epoch - 14ms/step
Epoch 89/100
20/20 - 0s - loss: 7.3972 - val_loss: 7.5219 - lr: 0.0010 - 283ms/epoch - 14ms/step
Epoch 90/100
20/20 - 0s - loss: 7.3960 - val_loss: 7.5206 - lr: 0.0010 - 279ms/epoch - 14ms/step
Epoch 91/100
20/20 - 0s - loss: 7.3914 - val_loss: 7.5241 - lr: 0.0010 - 310ms/epoch - 16ms/step
Epoch 92/100
20/20 - 0s - loss: 7.3944 - val_loss: 7.5266 - lr: 0.0010 - 282ms/epoch - 14ms/step
Epoch 93/100
20/20 - 0s - loss: 7.3969 - val_loss: 7.5102 - lr: 0.0010 - 289ms/epoch - 14ms/step
Epoch 94/100
20/20 - 0s - loss: 7.3884 - val_loss: 7.5156 - lr: 0.0010 - 287ms/epoch - 14ms/step
Epoch 95/100
20/20 - 0s - loss: 7.3845 - val_loss: 7.5140 - lr: 0.0010 - 286ms/epoch - 14ms/step
Epoch 96/100
20/20 - 0s - loss: 7.3824 - val_loss: 7.5099 - lr: 0.0010 - 280ms/epoch - 14ms/step
Epoch 97/100
20/20 - 0s - loss: 7.3811 - val_loss: 7.5223 - lr: 0.0010 - 285ms/epoch - 14ms/step
Epoch 98/100
20/20 - 0s - loss: 7.3790 - val_loss: 7.5190 - lr: 0.0010 - 293ms/epoch - 15ms/step
Epoch 99/100
20/20 - 0s - loss: 7.3801 - val_loss: 7.5113 - lr: 0.0010 - 280ms/epoch - 14ms/step
Epoch 100/100
20/20 - 0s - loss: 7.3802 - val_loss: 7.5080 - lr: 0.0010 - 288ms/epoch - 14ms/step
* Population optimizer:
- best model is number 1
- best loss function is 7.38
- best validation loss function is 7.51
- population losses [7.51]
# we can plot training summaries to make sure training went smoothly:
flow.training_plot()
plt.show();
# and we can print the training summary:
flow.print_training_summary()
loss : 7.3802 val_loss : 7.5080 lr : 0.0010 chi2Z_ks : 0.0269 chi2Z_ks_p : 0.4554 loss_rate : 6.5804e-05 val_loss_rate: -0.0034
# we can triangle plot the flow to see how well it has learned the target distribution:
g = plots.get_subplot_plotter()
g.triangle_plot([chain, flow.MCSamples(20000)],
params=flow.param_names,
filled=True)
plt.show();
# this looks nice but not perfect, let's train for longer:
flow.feedback = 1
flow.train(epochs=300, verbose=-1); # verbose = -1 uses tqdm progress bar
0epoch [00:00, ?epoch/s]
0batch [00:00, ?batch/s]
# we can plot training summaries to make sure training went smoothly:
flow.training_plot()
plt.show();
<Figure size 640x480 with 0 Axes>
If you train for long enough you should start seeing the learning rate adapting to the non-improving (noisy) loss function.
This means that the flow is learning finer and finer features and a good indication that training is converging. If you push it further, at some point, the flow will start overfitting and training will stop.
Now let's look at how the marginal distributions look like:
# we can triangle plot the flow to see how well it has learned the target distribution:
g = plots.get_subplot_plotter()
g.triangle_plot([chain,
flow.MCSamples(20000) # this flow method returns a MCSamples object
],
params=flow.param_names,
filled=True)
plt.show();
This is now much better!
We can use the trained flow to perform several operations. For example let's compute log-likelihoods
samples = flow.MCSamples(20000)
logP = flow.log_probability(flow.cast(samples.samples)).numpy()
samples.addDerived(logP, name='logP', label='\\log P')
samples.updateBaseStatistics();
# now let's plot everything:
g = plots.get_subplot_plotter()
g.triangle_plot([samples, chain],
plot_3d_with_param='logP',
filled=False)
plt.show();
We can appreciate here a beautiful display of a projection effect. The marginal distribution of $p_5$ is peaked at a positive value while the logP plot clearly shows that the peak of the full distribution is the negative one.
If you are interested in understanding systematically these types of effect, check the corresponding tensiometer tutorial!
Average flow example:¶
A more advanced flow model consists in training several flows and using a weighted mixture normalizing flow model.
This flow model improves the variance of the flow in regions that are scarse with samples (as different flow models will allucinate differently)...
Let's try averaging 5 flow models (note that we could do this in parallel with MPI on bigger machines):
kwargs = {
'feedback': 1,
'verbose': -1,
'plot_every': 1000,
'pop_size': 1,
'num_flows': 5,
'epochs': 400,
}
average_flow = sp.average_flow_from_chain(chain, # parameter difference chain
**kwargs)
Training flow 0
0epoch [00:00, ?epoch/s]
Training flow 1
0epoch [00:00, ?epoch/s]
Training flow 2
0epoch [00:00, ?epoch/s]
Training flow 3
0epoch [00:00, ?epoch/s]
Training flow 4
0epoch [00:00, ?epoch/s]
# most methods are implemented for the average flow as well:
average_flow.training_plot()
plt.show();