RelaxedBernoulli
DistributionRelaxedBernoulli(temperature: Tensor | float, probs: Tensor | float | None = None, logits: Tensor | float | None = None, validate_args: bool | None = None)Concrete (Gumbel-sigmoid) relaxation of the Bernoulli distribution.
RelaxedBernoulli(temperature=τ, probs=p) defines a continuous
distribution over whose samples are differentiable
surrogates for Bernoulli samples. As the distribution
concentrates on , recovering the discrete Bernoulli.
As the distribution approaches
.
The Concrete distribution (Maddison et al. 2017) / Gumbel-softmax (Jang et al. 2017) trick enables gradient-based optimisation through discrete latent variables in variational autoencoders and related models.
Parameters
temperatureTensor | floatlogits.probs.validate_argsbool | None= NoneTrue, validate parameter constraints at construction time.Attributes
Notes
Reparameterised sampling (Gumbel-sigmoid trick):
where is the sigmoid function. Gradients propagate through both and .
Log-PDF (Maddison et al. 2017, Eq. 2):
where .
Examples
>>> import lucid
>>> from lucid.distributions import RelaxedBernoulli
>>> lucid.manual_seed(0)
>>> dist = RelaxedBernoulli(temperature=0.5, probs=0.7)
>>> samples = dist.rsample((100,))
>>> # Samples are in (0, 1)
>>> bool(((samples > 0) & (samples < 1)).all())
TrueUsed by 1
Constructors
1__init__
→None__init__(temperature: Tensor | float, probs: Tensor | float | None = None, logits: Tensor | float | None = None, validate_args: bool | None = None)Construct a RelaxedBernoulli (Binary Concrete) distribution.
Parameters
temperatureTensor | floatlogits.probs.validate_argsbool | None= NoneTrue, validate parameter constraints at construction time.Raises
ValueErrorprobs and logits are provided.Instance methods
4Log-odds — as given, or derived on access.
Derived from probs clamped one epsilon inside , as
the reference framework derives it.
Success probability — as given, or sigmoid(logits) on access.
Draw a reparameterised sample via the Gumbel-sigmoid trick.
Parameters
sample_shapetuple[int, ...]= ()Returns
TensorSamples of shape
(*sample_shape, *batch_shape), with gradients flowing through
both logits and temperature.