Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
43 changes: 42 additions & 1 deletion tests/test_flows.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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)

Expand Down
7 changes: 7 additions & 0 deletions tests/test_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 2 additions & 0 deletions zuko/flows/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
129 changes: 129 additions & 0 deletions zuko/flows/planar.py
Original file line number Diff line number Diff line change
@@ -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=<SubBackward0>)
"""

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)
Loading