diff --git a/nflows/transforms/splines/rational_quadratic.py b/nflows/transforms/splines/rational_quadratic.py index bfe57f3..352b3fe 100644 --- a/nflows/transforms/splines/rational_quadratic.py +++ b/nflows/transforms/splines/rational_quadratic.py @@ -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 @@ -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 @@ -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] @@ -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) diff --git a/tests/transforms/splines/rational_quadratic_test.py b/tests/transforms/splines/rational_quadratic_test.py index 95308b7..1eb7b8e 100644 --- a/tests/transforms/splines/rational_quadratic_test.py +++ b/tests/transforms/splines/rational_quadratic_test.py @@ -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): @@ -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])