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
98 changes: 61 additions & 37 deletions nflows/transforms/splines/rational_quadratic.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@

from ...transforms.base import InputOutsideDomain
from ...utils import torchutils
from .utils import apply_spline_at_mask

DEFAULT_MIN_BIN_WIDTH = 1e-3
DEFAULT_MIN_BIN_HEIGHT = 1e-3
Expand All @@ -23,43 +22,58 @@ def unconstrained_rational_quadratic_spline(
min_bin_height=DEFAULT_MIN_BIN_HEIGHT,
min_derivative=DEFAULT_MIN_DERIVATIVE,
):
inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound)
outside_interval_mask = ~inside_interval_mask

outputs = torch.zeros_like(inputs)
logabsdet = torch.zeros_like(inputs)

if tails == "linear":
unnormalized_derivatives = F.pad(unnormalized_derivatives, pad=(1, 1))
constant = np.log(np.exp(1 - min_derivative) - 1)
unnormalized_derivatives[..., 0] = constant
unnormalized_derivatives[..., -1] = constant

outputs[outside_interval_mask] = inputs[outside_interval_mask]
logabsdet[outside_interval_mask] = 0
else:
"""Rational-quadratic spline with linear tails outside [-tail_bound, tail_bound].

Implemented with static shapes: the spline is evaluated on the whole tensor
(inputs clamped into the domain) and the tails are selected with
``torch.where``. Compared to masked indexing (``outputs[mask] = ...``) this
involves no data-dependent shapes and no host synchronisation, so the
function can be captured by ``torch.compile`` without graph breaks and its
many small kernels can be fused. The result is numerically identical to the
masked formulation.
"""
if tails != "linear":
raise RuntimeError("{} tails are not implemented.".format(tails))

if torch.any(inside_interval_mask):
outputs, logabsdet = apply_spline_at_mask(
rational_quadratic_spline,
inputs=inputs,
outputs=outputs,
logabsdet=logabsdet,
mask=inside_interval_mask,
unnormalized_widths=unnormalized_widths[inside_interval_mask, :],
unnormalized_heights=unnormalized_heights[inside_interval_mask, :],
unnormalized_derivatives=unnormalized_derivatives[inside_interval_mask, :],
inverse=inverse,
left=-tail_bound,
right=tail_bound,
bottom=-tail_bound,
top=tail_bound,
min_bin_width=min_bin_width,
min_bin_height=min_bin_height,
min_derivative=min_derivative,
inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound)

)
unnormalized_derivatives = F.pad(unnormalized_derivatives, pad=(1, 1))
constant = np.log(np.exp(1 - min_derivative) - 1)
unnormalized_derivatives[..., 0] = constant
unnormalized_derivatives[..., -1] = constant

# Inputs in the tails are clamped into the spline domain so the spline is
# finite everywhere (their spline value is discarded below); inputs inside
# the domain, bounds included, pass through unchanged so that their gradient
# is exactly that of the masked formulation (the subgradient of clamp at its
# bounds differs between torch versions). With every input inside the
# domain, the (host-syncing) domain check of the inner spline is redundant.
spline_inputs = torch.where(
inside_interval_mask, inputs, inputs.clamp(-tail_bound, tail_bound)
)
spline_outputs, spline_logabsdet = rational_quadratic_spline(
inputs=spline_inputs,
unnormalized_widths=unnormalized_widths,
unnormalized_heights=unnormalized_heights,
unnormalized_derivatives=unnormalized_derivatives,
inverse=inverse,
left=-tail_bound,
right=tail_bound,
bottom=-tail_bound,
top=tail_bound,
min_bin_width=min_bin_width,
min_bin_height=min_bin_height,
min_derivative=min_derivative,
check_domain=False,
)
outputs = torch.where(
inside_interval_mask, spline_outputs.to(inputs.dtype), inputs
)
logabsdet = torch.where(
inside_interval_mask,
spline_logabsdet.to(inputs.dtype),
torch.zeros_like(inputs),
)
return outputs, logabsdet


Expand All @@ -76,8 +90,17 @@ def rational_quadratic_spline(
min_bin_width=DEFAULT_MIN_BIN_WIDTH,
min_bin_height=DEFAULT_MIN_BIN_HEIGHT,
min_derivative=DEFAULT_MIN_DERIVATIVE,
check_domain=True,
):
if torch.min(inputs) < left or torch.max(inputs) > right:
"""Rational-quadratic spline on [left, right] x [bottom, top].

``check_domain=False`` skips the host-synchronising checks that the inputs
lie in the domain (and, for the inverse, that the discriminant is
non-negative). Only use it when the caller guarantees the domain, e.g. after
clamping; it lets the function be captured by ``torch.compile`` without a
graph break.
"""
if check_domain and (torch.min(inputs) < left or torch.max(inputs) > right):
raise InputOutsideDomain()

num_bins = unnormalized_widths.shape[-1]
Expand Down Expand Up @@ -134,7 +157,8 @@ def rational_quadratic_spline(
c = -input_delta * (inputs - input_cumheights)

discriminant = b.pow(2) - 4 * a * c
assert (discriminant >= 0).all()
if check_domain:
assert (discriminant >= 0).all()

root = (2 * c) / (-b - torch.sqrt(discriminant))
# root = (- b + torch.sqrt(discriminant)) / (2 * a)
Expand Down
106 changes: 106 additions & 0 deletions tests/transforms/splines/rational_quadratic_test.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
import numpy as np
import torch
import torchtestcase
from torch.nn import functional as F

from nflows.transforms import splines
from nflows.transforms.splines import rational_quadratic as rq


class RationalQuadraticSplineTest(torchtestcase.TorchTestCase):
Expand Down Expand Up @@ -82,3 +85,106 @@ def call_spline_fn(inputs, inverse=False):
self.eps = 1e-4
self.assertEqual(inputs, inputs_inv)
self.assertEqual(logabsdet + logabsdet_inv, torch.zeros_like(logabsdet))


def _reference_unconstrained_rational_quadratic_spline(
inputs,
unnormalized_widths,
unnormalized_heights,
unnormalized_derivatives,
inverse=False,
tail_bound=1.0,
):
"""The previous (masked-indexing) implementation, kept as a numerical reference."""
inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound)
outside_interval_mask = ~inside_interval_mask
outputs = torch.zeros_like(inputs)
logabsdet = torch.zeros_like(inputs)
unnormalized_derivatives = F.pad(unnormalized_derivatives, pad=(1, 1))
constant = np.log(np.exp(1 - rq.DEFAULT_MIN_DERIVATIVE) - 1)
unnormalized_derivatives[..., 0] = constant
unnormalized_derivatives[..., -1] = constant
outputs[outside_interval_mask] = inputs[outside_interval_mask]
logabsdet[outside_interval_mask] = 0
if torch.any(inside_interval_mask):
outputs_masked, logabsdet_masked = rq.rational_quadratic_spline(
inputs=inputs[inside_interval_mask],
unnormalized_widths=unnormalized_widths[inside_interval_mask, :],
unnormalized_heights=unnormalized_heights[inside_interval_mask, :],
unnormalized_derivatives=unnormalized_derivatives[inside_interval_mask, :],
inverse=inverse,
left=-tail_bound,
right=tail_bound,
bottom=-tail_bound,
top=tail_bound,
)
outputs[inside_interval_mask] = outputs_masked
logabsdet[inside_interval_mask] = logabsdet_masked
return outputs, logabsdet


class UnconstrainedRationalQuadraticSplineStaticShapeTest(torchtestcase.TorchTestCase):
"""The unconstrained spline is written with static shapes (clamp + where) so it
can be captured by torch.compile. It must remain bit-identical to the masked
reference implementation, including gradients."""

def _make_inputs(self, dtype, mode):
torch.manual_seed(1234)
num_bins = 8
shape = [16, 5]
if mode == "mixed":
inputs = 3 * torch.randn(*shape, dtype=dtype)
elif mode == "all_outside":
inputs = torch.sign(torch.randn(*shape)) * (1.0 + torch.rand(*shape))
inputs = inputs.to(dtype)
elif mode == "boundary": # exactly on the tail bounds
inputs = torch.tensor([-1.0, 1.0, 0.0, -1.0, 1.0] * 16, dtype=dtype).reshape(shape)
params = [
torch.randn(*shape, num_bins, dtype=dtype),
torch.randn(*shape, num_bins, dtype=dtype),
torch.randn(*shape, num_bins - 1, dtype=dtype),
]
return inputs, params

def _check(self, dtype, mode, inverse):
inputs, params = self._make_inputs(dtype, mode)

def run(fn):
x = inputs.clone().requires_grad_(True)
ps = [p.clone().requires_grad_(True) for p in params]
out, logabsdet = fn(x, *ps, inverse=inverse)
(out.sum() + logabsdet.sum()).backward()
return out, logabsdet, [x.grad] + [p.grad for p in ps]

ref = run(_reference_unconstrained_rational_quadratic_spline)
new = run(splines.unconstrained_rational_quadratic_spline)
self.assertEqual(ref[0].dtype, new[0].dtype)
self.assertTrue(torch.equal(ref[0], new[0]), "outputs differ")
self.assertTrue(torch.equal(ref[1], new[1]), "logabsdet differ")
for g_ref, g_new in zip(ref[2], new[2]):
# all_outside: the reference never touches the spline parameters (grad
# None); the static-shape version gives an explicit zero, equivalent.
if g_ref is None:
self.assertTrue(g_new is None or not g_new.any(), "gradients differ")
else:
self.assertTrue(torch.equal(g_ref, g_new), "gradients differ")

def test_matches_reference_implementation(self):
for dtype in (torch.float32, torch.float64):
for mode in ("mixed", "all_outside", "boundary"):
for inverse in (False, True):
with self.subTest(dtype=dtype, mode=mode, inverse=inverse):
self._check(dtype, mode, inverse)

def test_compiles_without_graph_breaks(self):
inputs, params = self._make_inputs(torch.float32, "mixed")
eager = splines.unconstrained_rational_quadratic_spline(inputs, *params)
compiled = torch.compile(
splines.unconstrained_rational_quadratic_spline,
fullgraph=True,
backend="aot_eager",
)
out, logabsdet = compiled(inputs, *params)
self.eps = 1e-6
self.assertEqual(out, eager[0])
self.assertEqual(logabsdet, eager[1])
Loading