diff --git a/README.md b/README.md index 5f2df28..d0c8779 100644 --- a/README.md +++ b/README.md @@ -93,6 +93,10 @@ For more information, check out the documentation and tutorials at [zuko.readthe | `CNF` | 2018 | [Neural Ordinary Differential Equations](https://arxiv.org/abs/1806.07366) | | `GF` | 2020 | [Gaussianization Flows](https://arxiv.org/abs/2003.01941) | | `BPF` | 2020 | [Bernstein-Polynomial Normalizing Flows](https://arxiv.org/abs/2004.00464) | +| `PNF` | 2015 | [Variational Inference with Normalizing Flows](https://arxiv.org/abs/1505.05770) | +| `TSNF` | 2018 | [Sylvester Normalizing Flows for Variational Inference](https://arxiv.org/abs/1803.05649) | +| `OSNF` | 2018 | [Sylvester Normalizing Flows for Variational Inference](https://arxiv.org/abs/1803.05649) | +| `HSNF` | 2018 | [Sylvester Normalizing Flows for Variational Inference](https://arxiv.org/abs/1803.05649) | ## Contributing diff --git a/tests/test_flows.py b/tests/test_flows.py index 60fdbf2..24037ac 100644 --- a/tests/test_flows.py +++ b/tests/test_flows.py @@ -10,7 +10,27 @@ from zuko.flows import * -@pytest.mark.parametrize("F", [NICE, MAF, NSF, SOSPF, NAF, UNAF, CNF, GF, BPF]) +@pytest.mark.parametrize( + "F", + [ + NICE, + MAF, + NSF, + SOSPF, + NAF, + UNAF, + CNF, + GF, + BPF, + PNF, + TSNF, + partial(TSNF, randperm=True), + OSNF, + partial(OSNF, hidden=2), + HSNF, + partial(HSNF, hidden=2, reflections=4), + ], +) def test_flows(tmp_path: Path, F: type) -> None: flow = F(3, 5) @@ -94,6 +114,27 @@ def test_flows(tmp_path: Path, F: type) -> None: assert repr(flow) +@pytest.mark.parametrize("F", [PNF, TSNF, OSNF, HSNF]) +def test_reversed_flows(F: type) -> None: + flow = F(3, 5) + flow = Flow(flow.transform.inv, flow.base) + + # Sampling and evaluation + c = randn(5) + x, log_p = flow(c).rsample_and_log_prob((32,)) + + assert x.shape == (32, 3) + assert torch.allclose(flow(c).log_prob(x), log_p, atol=1e-4) + + # Reparameterization trick + flow.zero_grad(set_to_none=True) + loss = (x.square().sum(dim=-1) + log_p).mean() + loss.backward() + + for name, p in flow.named_parameters(): + assert p.grad is not None, name + + def test_triangular_transforms() -> None: order = torch.randperm(5) diff --git a/tests/test_transforms.py b/tests/test_transforms.py index d639d44..95299c9 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -83,6 +83,13 @@ def test_multivariate_transforms() -> None: PermutationTransform(torch.randperm(5)), RotationTransform(randn(5, 5)), LULinearTransform(randn(5, 5)), + PlanarTransform(5 * randn(5), randn(5), randn(())), + SylvesterTransform(torch.linalg.qr(randn(5, 3)).Q, randn(3, 3), randn(3, 3), randn(3)), + SylvesterTransform( + torch.linalg.qr(randn(5, 5)).Q, 3 * randn(5, 5), 3 * randn(5, 5), randn(5) + ), + TriangularSylvesterTransform(randn(5, 5), randn(5, 5), randn(5)), + TriangularSylvesterTransform(randn(5, 5), randn(5, 5), randn(5), torch.randperm(5)), ] for t in ts: diff --git a/zuko/flows/__init__.py b/zuko/flows/__init__.py index 32ef699..fd2d38b 100644 --- a/zuko/flows/__init__.py +++ b/zuko/flows/__init__.py @@ -14,5 +14,7 @@ from .gaussianization import GF, ElementWiseTransform from .mixture import GMM from .neural import NAF, UNAF +from .planar import PNF, PlanarLazyTransform from .polynomial import BPF, SOSPF from .spline import NCSF, NSF +from .sylvester import HSNF, OSNF, TSNF, SylvesterLazyTransform diff --git a/zuko/flows/planar.py b/zuko/flows/planar.py new file mode 100644 index 0000000..d9b36d4 --- /dev/null +++ b/zuko/flows/planar.py @@ -0,0 +1,129 @@ +r"""Planar flows.""" + +__all__ = [ + "PNF", + "PlanarLazyTransform", +] + +import torch +import torch.nn as nn + +from math import prod +from torch import Tensor +from torch.distributions import Transform + +from ..distributions import DiagNormal +from ..lazy import Flow, LazyTransform, UnconditionalDistribution +from ..nn import MLP +from ..transforms import PlanarTransform +from ..utils import unpack + + +class PlanarLazyTransform(LazyTransform): + r"""Creates a lazy planar transformation. + + See also: + :class:`zuko.transforms.PlanarTransform` + + References: + | Variational Inference with Normalizing Flows (Rezende et al., 2015) + | https://arxiv.org/abs/1505.05770 + + Arguments: + features: The number of features. + context: The number of context features. + kwargs: Keyword arguments passed to :class:`zuko.nn.MLP`. + + Example: + >>> t = PlanarLazyTransform(3, 4) + >>> t + PlanarLazyTransform( + (hyper): MLP( + (0): Linear(in_features=4, out_features=64, bias=True) + (1): ReLU() + (2): Linear(in_features=64, out_features=64, bias=True) + (3): ReLU() + (4): Linear(in_features=64, out_features=7, bias=True) + ) + ) + >>> x = torch.randn(3) + >>> x + tensor([-1.3411, 0.1149, 0.6243]) + >>> c = torch.randn(4) + >>> y = t(c)(x) + >>> t(c).inv(y) + tensor([-1.3411, 0.1149, 0.6243], grad_fn=) + """ + + def __init__( + self, + features: int, + context: int = 0, + **kwargs, + ) -> None: + super().__init__() + + self.shapes = [(features,), (features,), ()] + self.total = sum(prod(s) for s in self.shapes) + + if context > 0: + self.hyper = MLP(context, self.total, **kwargs) + else: + self.phi = nn.ParameterList(torch.randn(s) for s in self.shapes) + + def forward(self, c: Tensor | None = None) -> Transform: + if c is None: + phi = self.phi + else: + phi = unpack(self.hyper(c), self.shapes) + + return PlanarTransform(*phi) + + +class PNF(Flow): + r"""Creates a planar normalizing flow (PNF). + + Note: + Evaluating densities is cheap, but sampling requires inverting planar + transformations with the bisection method. When sampling is the bottleneck, as + in variational inference, the flow can be reversed with + :py:`Flow(flow.transform.inv, flow.base)`, as in the original paper. + + See also: + :class:`PlanarLazyTransform` + + References: + | Variational Inference with Normalizing Flows (Rezende et al., 2015) + | https://arxiv.org/abs/1505.05770 + + Arguments: + features: The number of features. + context: The number of context features. + transforms: The number of planar transformations. + kwargs: Keyword arguments passed to :class:`PlanarLazyTransform`. + """ + + def __init__( + self, + features: int, + context: int = 0, + transforms: int = 3, + **kwargs, + ) -> None: + transforms = [ + PlanarLazyTransform( + features=features, + context=context, + **kwargs, + ) + for _ in range(transforms) + ] + + base = UnconditionalDistribution( + DiagNormal, + loc=torch.zeros(features), + scale=torch.ones(features), + buffer=True, + ) + + super().__init__(transforms, base) diff --git a/zuko/flows/sylvester.py b/zuko/flows/sylvester.py new file mode 100644 index 0000000..48a8165 --- /dev/null +++ b/zuko/flows/sylvester.py @@ -0,0 +1,335 @@ +r"""Sylvester flows.""" + +__all__ = [ + "HSNF", + "OSNF", + "TSNF", + "SylvesterLazyTransform", +] + +import math +import torch +import torch.nn as nn + +from math import prod +from torch import LongTensor, Tensor +from torch.distributions import Transform +from typing import Literal + +from ..distributions import DiagNormal +from ..lazy import Flow, LazyTransform, UnconditionalDistribution +from ..nn import MLP +from ..transforms import SylvesterTransform, TriangularSylvesterTransform +from ..utils import unpack + + +class SylvesterLazyTransform(LazyTransform): + r"""Creates a lazy Sylvester transformation. + + The upper triangular matrices :math:`R_1` and :math:`R_2` are parameterized by + their :math:`M (M + 1) / 2` upper triangular elements. The matrix :math:`Q` with + orthonormal columns is parameterized according to the :py:`orthogonal` strategy: + + * :py:`'householder'`: the first :math:`M` columns of a product of :math:`K` + Householder reflections. + * :py:`'qr'`: the orthonormal factor of the QR decomposition of an unconstrained + :math:`D \times M` matrix. + * :py:`'permutation'`: a fixed permutation matrix, which requires :math:`M = D`. + + See also: + | :class:`zuko.transforms.SylvesterTransform` + | :class:`zuko.transforms.TriangularSylvesterTransform` + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + features: The number of features. + context: The number of context features. + hidden: The number of hidden units :math:`M`. If :py:`None`, :math:`M = D`. + orthogonal: The parameterization of :math:`Q`. One of + :py:`['householder', 'qr', 'permutation']`. + reflections: The number of Householder reflections :math:`K`. If :py:`None`, + :math:`K = M`. Only used when :py:`orthogonal='householder'`. + order: The feature permutation, with shape :math:`(D,)`. If :py:`None`, the + identity is used. Only used when :py:`orthogonal='permutation'`. + kwargs: Keyword arguments passed to :class:`zuko.nn.MLP`. + + Example: + >>> t = SylvesterLazyTransform(3, 4) + >>> t + SylvesterLazyTransform( + (orthogonal): householder + (hyper): MLP( + (0): Linear(in_features=4, out_features=64, bias=True) + (1): ReLU() + (2): Linear(in_features=64, out_features=64, bias=True) + (3): ReLU() + (4): Linear(in_features=64, out_features=24, bias=True) + ) + ) + >>> x = torch.randn(3) + >>> x + tensor([ 0.1991, 0.5514, -1.4672]) + >>> c = torch.randn(4) + >>> y = t(c)(x) + >>> t(c).inv(y) + tensor([ 0.1991, 0.5514, -1.4672], grad_fn=) + """ + + def __init__( + self, + features: int, + context: int = 0, + hidden: int | None = None, + orthogonal: Literal["householder", "qr", "permutation"] = "householder", + reflections: int | None = None, + order: LongTensor | None = None, + **kwargs, + ) -> None: + super().__init__() + + if hidden is None: + hidden = features + + assert hidden <= features, "'hidden' should not be greater than 'features'" + + if orthogonal == "householder": + if reflections is None: + reflections = hidden + + shapes = [(reflections, features)] + elif orthogonal == "qr": + shapes = [(features, hidden)] + elif orthogonal == "permutation": + assert hidden == features, "'hidden' should equal 'features' for permutations" + + shapes = [] + else: + raise ValueError(f"unknown orthogonal parameterization '{orthogonal}'") + + triu = hidden * (hidden + 1) // 2 + + self.hidden = hidden + self.orthogonal = orthogonal + self.register_buffer("order", order) + + self.shapes = [*shapes, (triu,), (triu,), (hidden,)] + self.total = sum(prod(s) for s in self.shapes) + + if context > 0: + self.hyper = MLP(context, self.total, **kwargs) + else: + self.phi = nn.ParameterList(torch.randn(s) / math.sqrt(s[-1]) for s in self.shapes) + + def extra_repr(self) -> str: + lines = [f"(orthogonal): {self.orthogonal}"] + + if self.order is not None: + lines.append(f"(order): {self.order.tolist()}") + + return "\n".join(lines) + + def triangular(self, r: Tensor) -> Tensor: + r"""Scatters the upper triangular elements into an :math:`M \times M` matrix.""" + + i, j = torch.triu_indices(self.hidden, self.hidden, device=r.device) + + R = r.new_zeros(*r.shape[:-1], self.hidden, self.hidden) + R[..., i, j] = r + + return R + + def householder(self, V: Tensor) -> Tensor: + r"""Returns the first :math:`M` columns of the product of the Householder + reflections defined by the rows of :math:`V`.""" + + V = V / torch.linalg.vector_norm(V, dim=-1, keepdim=True).clamp(min=1e-8) + Q = torch.eye(V.shape[-1], dtype=V.dtype, device=V.device)[:, : self.hidden] + + for v in torch.unbind(V, dim=-2): + Q = Q - 2 * v.unsqueeze(-1) * (v.unsqueeze(-2) @ Q) + + return Q + + def forward(self, c: Tensor | None = None) -> Transform: + if c is None: + phi = self.phi + else: + phi = unpack(self.hyper(c), self.shapes) + + *phi, r1, r2, b = phi + R1, R2 = self.triangular(r1), self.triangular(r2) + + if self.orthogonal == "householder": + Q = self.householder(*phi) + elif self.orthogonal == "qr": + Q, _ = torch.linalg.qr(*phi) + else: + return TriangularSylvesterTransform(R1, R2, b, order=self.order) + + return SylvesterTransform(Q, R1, R2, b) + + +class TSNF(Flow): + r"""Creates a triangular Sylvester normalizing flow (TSNF). + + Note: + Evaluating densities is cheap, but sampling requires inverting Sylvester + transformations, which involves :math:`M` sequential root-finding problems per + transformation. When sampling is the bottleneck, as in variational inference, + the flow can be reversed with :py:`Flow(flow.transform.inv, flow.base)`, which + makes sampling cheap and density evaluation expensive, as in the original paper. + + See also: + :class:`SylvesterLazyTransform` + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + features: The number of features. + context: The number of context features. + transforms: The number of Sylvester transformations. + randperm: Whether the features are randomly permuted between transformations + or not. If :py:`False`, features are in ascending (descending) order for + even (odd) transformations. + kwargs: Keyword arguments passed to :class:`SylvesterLazyTransform`. + """ + + def __init__( + self, + features: int, + context: int = 0, + transforms: int = 3, + randperm: bool = False, + **kwargs, + ) -> None: + orders = [ + torch.arange(features), + torch.flipud(torch.arange(features)), + ] + + transforms = [ + SylvesterLazyTransform( + features=features, + context=context, + orthogonal="permutation", + order=torch.randperm(features) if randperm else orders[i % 2], + **kwargs, + ) + for i in range(transforms) + ] + + base = UnconditionalDistribution( + DiagNormal, + loc=torch.zeros(features), + scale=torch.ones(features), + buffer=True, + ) + + super().__init__(transforms, base) + + +class OSNF(Flow): + r"""Creates an orthogonal Sylvester normalizing flow (OSNF). + + Note: + Evaluating densities is cheap, but sampling requires inverting Sylvester + transformations, which involves :math:`M` sequential root-finding problems per + transformation. When sampling is the bottleneck, as in variational inference, + the flow can be reversed with :py:`Flow(flow.transform.inv, flow.base)`, which + makes sampling cheap and density evaluation expensive, as in the original paper. + + See also: + :class:`SylvesterLazyTransform` + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + features: The number of features. + context: The number of context features. + transforms: The number of Sylvester transformations. + kwargs: Keyword arguments passed to :class:`SylvesterLazyTransform`. + """ + + def __init__( + self, + features: int, + context: int = 0, + transforms: int = 3, + **kwargs, + ) -> None: + transforms = [ + SylvesterLazyTransform( + features=features, + context=context, + orthogonal="qr", + **kwargs, + ) + for _ in range(transforms) + ] + + base = UnconditionalDistribution( + DiagNormal, + loc=torch.zeros(features), + scale=torch.ones(features), + buffer=True, + ) + + super().__init__(transforms, base) + + +class HSNF(Flow): + r"""Creates a Householder Sylvester normalizing flow (HSNF). + + Note: + Evaluating densities is cheap, but sampling requires inverting Sylvester + transformations, which involves :math:`M` sequential root-finding problems per + transformation. When sampling is the bottleneck, as in variational inference, + the flow can be reversed with :py:`Flow(flow.transform.inv, flow.base)`, which + makes sampling cheap and density evaluation expensive, as in the original paper. + + See also: + :class:`SylvesterLazyTransform` + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + features: The number of features. + context: The number of context features. + transforms: The number of Sylvester transformations. + kwargs: Keyword arguments passed to :class:`SylvesterLazyTransform`. + """ + + def __init__( + self, + features: int, + context: int = 0, + transforms: int = 3, + **kwargs, + ) -> None: + transforms = [ + SylvesterLazyTransform( + features=features, + context=context, + orthogonal="householder", + **kwargs, + ) + for _ in range(transforms) + ] + + base = UnconditionalDistribution( + DiagNormal, + loc=torch.zeros(features), + scale=torch.ones(features), + buffer=True, + ) + + super().__init__(transforms, base) diff --git a/zuko/transforms.py b/zuko/transforms.py index 6383314..fb42e39 100644 --- a/zuko/transforms.py +++ b/zuko/transforms.py @@ -18,11 +18,14 @@ "MonotonicRQSTransform", "MonotonicTransform", "PermutationTransform", + "PlanarTransform", "RotationTransform", "SOSPolynomialTransform", "SignedPowerTransform", "SinTransform", "SoftclipTransform", + "SylvesterTransform", + "TriangularSylvesterTransform", "UnconstrainedMonotonicTransform", ] @@ -1179,6 +1182,226 @@ def f_aug(t: Tensor, x: Tensor, ladj: Tensor) -> Tensor: return y, ladj * (1 / self.trace_scale) +class _TanhResidualTransform(Transform): + r"""Creates a transformation of the form + + .. math:: f(x) = x + A \tanh(B x + b) + + where :math:`A` is a :math:`D \times M` matrix, :math:`B` is a :math:`M \times D` + matrix and :math:`C = B A` is an upper triangular matrix whose diagonal elements + are greater than :math:`-1`. + + Arguments: + A: The matrix :math:`A`, with shape :math:`(*, D, M)`. + B: The matrix :math:`B`, with shape :math:`(*, M, D)`. + C: The upper triangular matrix :math:`C = B A`, with shape :math:`(*, M, M)`. + b: The bias vector :math:`b`, with shape :math:`(*, M)`. + """ + + domain = constraints.real_vector + codomain = constraints.real_vector + bijective = True + + def __init__(self, A: Tensor, B: Tensor, C: Tensor, b: Tensor, **kwargs) -> None: + super().__init__(**kwargs) + + self.A = A + self.B = B + self.C = C + self.b = b + + def _call(self, x: Tensor) -> Tensor: + z = torch.einsum("...ij,...j->...i", self.B, x) + self.b + + return x + torch.einsum("...ij,...j->...i", self.A, torch.tanh(z)) + + def _inverse(self, y: Tensor) -> Tensor: + # Since B x = B y - C tanh(z) and C is upper triangular, the pre-activations + # z = B x + b are recovered by back-substitution. Each step is a monotonic + # univariate root-finding problem z_i + c_ii tanh(z_i) = r_i whose solution + # lies within [r_i - |c_ii|, r_i + |c_ii|]. + + c = torch.einsum("...ij,...j->...i", self.B, y) + self.b + h = [] + + for i in reversed(range(c.shape[-1])): + r = c[..., i] + + if h: + r = r - torch.einsum( + "...j,...j", self.C[..., i, i + 1 :], torch.stack(h[::-1], dim=-1) + ) + + s = self.C[..., i, i] + z = bisection( + f=lambda z, s=s: z + s * torch.tanh(z), + y=r, + a=r - s.abs(), + b=r + s.abs(), + n=64, + phi=(s,) if s.requires_grad else (), + ) + + h.append(torch.tanh(z)) + + h = torch.stack(h[::-1], dim=-1) + + return y - torch.einsum("...ij,...j->...i", self.A, h) + + def log_abs_det_jacobian(self, x: Tensor, y: Tensor) -> Tensor: + _, ladj = self.call_and_ladj(x) + return ladj + + def call_and_ladj(self, x: Tensor) -> tuple[Tensor, Tensor]: + z = torch.einsum("...ij,...j->...i", self.B, x) + self.b + h = torch.tanh(z) + y = x + torch.einsum("...ij,...j->...i", self.A, h) + + diag = torch.diagonal(self.C, dim1=-2, dim2=-1) + ladj = torch.log1p((1 - h**2) * diag).sum(dim=-1) + + return y, ladj + + +class SylvesterTransform(_TanhResidualTransform): + r"""Creates a Sylvester transformation. + + .. math:: f(x) = x + Q R_1 \tanh(R_2 Q^T x + b) + + where :math:`Q` is a :math:`D \times M` matrix with orthonormal columns + (:math:`Q^T Q = I`) and :math:`R_1` and :math:`R_2` are :math:`M \times M` upper + triangular matrices. + + To ensure invertibility, the diagonal elements of :math:`R_1` and :math:`R_2` are + mapped to the interval :math:`(-1, 1)` with the hyperbolic tangent function, which + guarantees :math:`r^1_{ii} \, r^2_{ii} > -1`. Thanks to Sylvester's determinant + identity and the triangular structure of :math:`R_1` and :math:`R_2`, the + log-absolute-determinant of the Jacobian is + + .. math:: \log |\det J_f(x)| = + \sum_{i = 1}^{M} \log \left( 1 + \tanh'(z_i) \, r^1_{ii} \, r^2_{ii} \right) + + where :math:`z = R_2 Q^T x + b`. The inverse transformation does not have a closed + form. Because :math:`R_2 R_1` is upper triangular, it reduces to :math:`M` + sequential univariate root-finding problems, which are solved with the bisection + method. + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + Q: A matrix :math:`Q` with orthonormal columns, with shape :math:`(*, D, M)`. + R1: A matrix whose upper triangular part is :math:`R_1`, with shape + :math:`(*, M, M)`. + R2: A matrix whose upper triangular part is :math:`R_2`, with shape + :math:`(*, M, M)`. + b: The bias vector :math:`b`, with shape :math:`(*, M)`. + """ + + def __init__(self, Q: Tensor, R1: Tensor, R2: Tensor, b: Tensor, **kwargs) -> None: + R1, R2 = self.triangular(R1), self.triangular(R2) + + super().__init__(Q @ R1, R2 @ Q.mT, R2 @ R1, b, **kwargs) + + @staticmethod + def triangular(R: Tensor) -> Tensor: + r"""Maps a matrix to an upper triangular matrix whose diagonal elements lie in + the interval :math:`(-1, 1)`.""" + + diag = torch.tanh(torch.diagonal(R, dim1=-2, dim2=-1)) + + return torch.triu(R, diagonal=1) + torch.diag_embed(diag) + + +class TriangularSylvesterTransform(_TanhResidualTransform): + r"""Creates a triangular Sylvester transformation. + + A special case of the Sylvester transformation for which :math:`M = D` and + :math:`Q` is a permutation matrix, such that :math:`Q^T x = x_\sigma` for a + permutation :math:`\sigma` of the features. + + See also: + :class:`SylvesterTransform` + + References: + | Sylvester Normalizing Flows for Variational Inference (van den Berg et al., 2018) + | https://arxiv.org/abs/1803.05649 + + Arguments: + R1: A matrix whose upper triangular part is :math:`R_1`, with shape + :math:`(*, D, D)`. + R2: A matrix whose upper triangular part is :math:`R_2`, with shape + :math:`(*, D, D)`. + b: The bias vector :math:`b`, with shape :math:`(*, D)`. + order: The permutation :math:`\sigma`, with shape :math:`(D,)`. If + :py:`None`, the identity is used. + """ + + def __init__( + self, + R1: Tensor, + R2: Tensor, + b: Tensor, + order: LongTensor | None = None, + **kwargs, + ) -> None: + R1 = SylvesterTransform.triangular(R1) + R2 = SylvesterTransform.triangular(R2) + + # Q R1 and R2 Q^T are row and column permutations, respectively + if order is None: + A, B = R1, R2 + else: + inverse = torch.argsort(order) + A, B = R1[..., inverse, :], R2[..., :, inverse] + + super().__init__(A, B, R2 @ R1, b, **kwargs) + + +class PlanarTransform(_TanhResidualTransform): + r"""Creates a planar transformation. + + .. math:: f(x) = x + \hat{u} \tanh(w^T x + b) + + The planar transformation is a Sylvester-type transformation with a single hidden + unit (:math:`M = 1`), but without orthogonality constraints. Following Rezende et + al. (2015), invertibility is ensured by the reparameterization + + .. math:: \hat{u} = u + \left( \log(1 + \exp(w^T u)) - 1 - w^T u \right) + \frac{w}{\|w\|^2} + + which guarantees :math:`w^T \hat{u} > -1`. The inverse transformation is obtained + by solving a univariate root-finding problem with the bisection method. + + See also: + :class:`SylvesterTransform` + + References: + | Variational Inference with Normalizing Flows (Rezende et al., 2015) + | https://arxiv.org/abs/1505.05770 + + Arguments: + u: The update vector :math:`u`, with shape :math:`(*, D)`. + w: The weight vector :math:`w`, with shape :math:`(*, D)`. + b: The bias :math:`b`, with shape :math:`(*,)`. + """ + + def __init__(self, u: Tensor, w: Tensor, b: Tensor, **kwargs) -> None: + uw = torch.sum(u * w, dim=-1, keepdim=True) + ww = torch.sum(w * w, dim=-1, keepdim=True) + u = u + (F.softplus(uw) - 1 - uw) * w / torch.clamp(ww, min=1e-12) + uw = torch.sum(u * w, dim=-1, keepdim=True) + + super().__init__( + u.unsqueeze(-1), + w.unsqueeze(-2), + uw.unsqueeze(-1), + b.unsqueeze(-1), + **kwargs, + ) + + class PermutationTransform(Transform): r"""Creates a transformation that permutes the elements.