TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

TransformedDistribution

class tensorplay.distributions.TransformedDistribution(base_distribution: Distribution, transforms: Transform | list[Transform], validate_args: bool | None = None)[source]

Extension of the Distribution class, which applies a sequence of Transforms to a base distribution. Let f be the composition of transforms applied:

X ~ BaseDistribution
Y = f(X) ~ TransformedDistribution(BaseDistribution, f)
log p(Y) = log p(X) + log |det (dX/dY)|

Note that the .event_shape of a TransformedDistribution is the maximum shape of its base distribution and its transforms, since transforms can introduce correlations among events.

An example for the usage of TransformedDistribution would be:

# Building a Logistic Distribution
# X ~ Uniform(0, 1)
# f = a + b * logit(X)
# Y ~ f(X) ~ Logistic(a, b)
base_distribution = Uniform(0, 1)
transforms = [SigmoidTransform().inv, AffineTransform(loc=a, scale=b)]
logistic = TransformedDistribution(base_distribution, transforms)

For more examples, please look at the implementations of Gumbel, HalfCauchy, HalfNormal, LogNormal, Pareto, Weibull, RelaxedBernoulli and RelaxedOneHotCategorical

property batch_shape: Size

Returns the shape over which parameters are batched.

cdf(value)[source]

Computes the cumulative distribution function by inverting the transform(s) and computing the score of the base distribution.

entropy() → Tensor

Returns entropy of distribution, batched over batch_shape.

Returns:

Tensor of shape batch_shape.

enumerate_support(expand: bool = True) → Tensor

Returns tensor containing all values supported by a discrete distribution. The result will enumerate over dimension 0, so the shape of the result will be (cardinality,) + batch_shape + event_shape (where event_shape = () for univariate distributions).

Note that this enumerates over all batched tensors in lock-step [[0, 0], [1, 1], …]. With expand=False, enumeration happens along dim 0, but with the remaining batch dimensions being singleton dimensions, [[0], [1], ...

To iterate over the full Cartesian product use itertools.product(m.enumerate_support()).

Parameters:

expand (bool) – whether to expand the support over the batch dims to match the distribution’s batch_shape.

Returns:

Tensor iterating over dimension 0.

property event_shape: Size

Returns the shape of a single sample (without batching).

icdf(value)[source]

Computes the inverse cumulative distribution function using transform(s) and computing the score of the base distribution.

log_prob(value)[source]

Scores the sample by inverting the transform(s) and computing the score using the score of the base distribution and the log abs det jacobian.

property mean: Tensor

Returns the mean of the distribution.

property mode: Tensor

Returns the mode of the distribution.

perplexity() → Tensor

Returns perplexity of distribution, batched over batch_shape.

Returns:

Tensor of shape batch_shape.

rsample(sample_shape: Size | Sequence[int] = ()) → Tensor[source]

Generates a sample_shape shaped reparameterized sample or sample_shape shaped batch of reparameterized samples if the distribution parameters are batched. Samples first from base distribution and applies transform() for every transform in the list.

sample(sample_shape=())[source]

Generates a sample_shape shaped sample or sample_shape shaped batch of samples if the distribution parameters are batched. Samples first from base distribution and applies transform() for every transform in the list.

sample_n(n: int) → Tensor

Generates n samples or n batches of samples if the distribution parameters are batched.

static set_default_validate_args(value: bool) → None

Sets whether validation is enabled or disabled.

The default behavior mimics Python’s assert statement: validation is on by default, but is disabled if Python is run in optimized mode (via python -O). Validation may be expensive, so you may want to disable it once a model is working.

Parameters:

value (bool) – Whether to enable validation.

property stddev: Tensor

Returns the standard deviation of the distribution.

property variance: Tensor

Returns the variance of the distribution.

On this page

Ask DeepWiki