RealNVPForImageGeneration
ImageGenerationModelRealNVPForImageGeneration(config: RealNVPConfig)RealNVP with the bits/dim training loss and .generate().
forward(x) returns a NormalizingFlowOutput whose loss
is the mean negative log-likelihood in bits per dimension — the
quantity the paper reports, and the scale-free form that keeps the
gradient magnitude independent of image size.
generate(n_samples) draws and
returns in a single parallel pass, back in [0, 1]
pixel space.
Parameters
configRealNVPConfigAttributes
realnvpRealNVPModelencode / decode /
log_prob / bits_per_dim.Notes
Reference: Dinh, Sohl-Dickstein, and Bengio, "Density Estimation Using Real NVP", ICLR, 2017 (arXiv:1605.08803). Reported bits/dim: 3.49 (CIFAR-10), 4.28 / 3.98 (Imagenet 32 / 64), 2.72–3.08 (LSUN), 3.02 (CelebA).
Examples
>>> import lucid
>>> from lucid.models.generative.realnvp import (
... RealNVPConfig, RealNVPForImageGeneration,
... )
>>> cfg = RealNVPConfig(sample_size=8, num_scales=2, residual_blocks=1,
... base_dim=8)
>>> model = RealNVPForImageGeneration(cfg).eval()
>>> out = model(lucid.rand((2, 3, 8, 8)))
>>> out.loss.shape # scalar bits/dim
()
>>> model.generate(n_samples=3).samples.shape
(3, 3, 8, 8)Used by 2
Constructors
1Instance methods
2forward(x: Tensor)generate(n_samples: int = 1, temperature: float = 1.0, device: str | None = None)Sample images by inverting a prior draw.
Parameters
n_samplesint= 1temperaturefloat= 1.01 trade diversity for typicality — a standard trick for
flows, not part of the paper.devicestr= NoneReturns
GenerationOutputsamples of shape (n_samples, C, H, W) in pixel space —
the logit stage is undone on the way out. The squash is not
saturating, so values land just outside [0, 1] (within
±(1 - c) / 2c); clip before display.
Examples
>>> import lucid
>>> from lucid.models.generative.realnvp import (
... RealNVPConfig, RealNVPForImageGeneration,
... )
>>> cfg = RealNVPConfig(sample_size=8, num_scales=2, residual_blocks=1,
... base_dim=8)
>>> model = RealNVPForImageGeneration(cfg).eval()
>>> samples = model.generate(n_samples=3).samples
>>> samples.shape
(3, 3, 8, 8)
>>> c = cfg.data_constraint
>>> pad = (1 - c) / (2 * c)
>>> bool(((samples >= -pad) & (samples <= 1 + pad)).all().item())
True
>>> model.generate(temperature=0.0)
Traceback (most recent call last):
...
ValueError: temperature must be positive, got 0.0