Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

NLL training of a Gaussianization flow

The same rotation + marginal blocks as RBIG, but fit end-to-end by maximum likelihood — the negative-log-likelihood objective and its log-det anatomy

00 — NLL training of a Gaussianization flow

Part 3 built RBIG: rotation + marginal blocks fit greedily, each layer optimised once against the data in front of it and never revisited. A parametric Gaussianization flow uses the same architecture — a stack of rotations and learnable marginal transforms — but treats every block’s parameters as free and fits them jointly, end-to-end, by maximum likelihood — the defining recipe for normalizing flows Papamakarios et al. (2021). The objective is the negative log-likelihood, and it comes straight from the change-of-variables rule (Part 0 00):

log⁡pX(x)=log⁡pZ(Tθ(x))+log⁡∣det⁡JTθ(x)∣,L(θ)=−1N∑ilog⁡pX(xi).\log p_X(x) = \log p_Z\big(T_\theta(x)\big) + \log\big|\det J_{T_\theta}(x)\big|, \qquad \mathcal{L}(\theta) = -\tfrac{1}{N}\sum_i \log p_X(x_i).

What you will see

1. The NLL objective and its anatomy

A flow TθT_\theta maps data xx to a latent z=Tθ(x)z = T_\theta(x) with base distribution pZ=N(0,I)p_Z = \mathcal{N}(0,I). Its log-density at xx has two parts: the base log-density of where xx lands, log⁡pZ(z)\log p_Z(z), and the log-determinant log⁡∣det⁡JTθ(x)∣\log|\det J_{T_\theta}(x)| that accounts for how the map stretches volume. Training minimises the mean NLL. We build an (untrained) flow and confirm the decomposition against its log_prob.

mean base log p_Z(z) = -2.041
mean log|det J|      = -14.894
sum                  = -16.934
flow.log_prob mean   = -16.934   (match: True)
=> initial NLL = 16.934

The two terms sum exactly to flow.log_prob — that identity is the loss the optimiser will minimise. The base term rewards mapping data to high-density regions of N(0,I)\mathcal{N}(0,I); the log-det term prevents the cheap cheat of just shrinking everything toward the origin (it penalises volume contraction). NLL training balances the two.

2. Train end-to-end with optax

gauss_flows ships a convenience trainer (fit_gaussianization_flow), but a Gaussianization flow is just an equinox module, so we can train it with a hand-built optax loop and add the bells and whistles that make deep flows converge: gradient clipping (clip_by_global_norm, so a rare exploding batch cannot wreck the parameters) and a cyclic learning-rate schedule (cosine_onecycle_schedule — warm up, then anneal to near zero). The loss is the NLL of §1; gradients come from eqx.filter_value_and_grad.

NLL: 16.934 (init) -> 1.978 (trained), 3000 steps
latent z after training: mean = +0.026, std = 0.704
<Figure size 1500x420 with 3 Axes>

The learning rate warms up then cosine-anneals to near zero (centre); the NLL falls steeply during the high-LR phase and settles as the rate decays (left); and the trained flow maps the two crescents toward a single Gaussian blob (right — not perfectly isotropic at this budget, but well-Gaussianized). Gradient clipping keeps the early high-LR steps from diverging. The flow is now a full generative model: evaluate log_prob for density, or sample the base and push through Tθ−1T_\theta^{-1} to generate.

3. Iterative vs parametric

Same architecture, two fitting philosophies. Greedy RBIG (fit_rbig, Part 3) fits each layer once, in sequence — fast, no gradients, no joint optimisation. Parametric training tunes all layers together against the likelihood. We fit both on two-moons and compare the learned densities.

mean log p(x):  greedy RBIG = -1.957   parametric = -1.978
<Figure size 1100x480 with 2 Axes>

Both recover the two-crescent density, and with enough layers and a tuned optax schedule the parametric flow matches the greedy RBIG fit on held-out likelihood. The difference is how they get there: greedy RBIG needs no gradients and is essentially instant, while the parametric flow pays for thousands of gradient steps from a random start. The natural question — can we have both, RBIG’s data-driven head start and gradient fine-tuning? — is exactly the next notebook. A greedy RBIG fit is an excellent initialisation, and 01 — RBIG warm-start shows warm-starting the trainable flow from RBIG converges far faster than the random start here.

Recap

piecerole
log⁡pX(x)=log⁡pZ(z)+log⁡∣det⁡J∣\log p_X(x) = \log p_Z(z) + \log\lvert\det J\rvertchange-of-variables density
NLL =−1N∑ilog⁡pX(xi)= -\frac1N\sum_i \log p_X(x_i)the training objective
base termrewards mapping data into high-density N(0,I)\mathcal{N}(0,I) regions
log-det termpenalises volume contraction (no shrink-to-origin cheat)
optax loopNLL + clip_by_global_norm + cyclic one-cycle cosine LR
greedy vs parametricgreedy = fit-once per layer; parametric = joint NLL (matches greedy when well-trained)

Next up. 01 — RBIG warm-start: initialise the trainable flow from a greedy RBIG fit and fine-tune — far faster convergence and a better optimum than training from a random start.

References
  1. Papamakarios, G., Nalisnick, E., Rezende, D. J., Mohamed, S., & Lakshminarayanan, B. (2021). Normalizing Flows for Probabilistic Modeling and Inference. Journal of Machine Learning Research, 22(57), 1–64.