From 63658c24ea89a6b840e44f2f1f38f58601dc9254 Mon Sep 17 00:00:00 2001 From: Nihar Gupte Date: Wed, 2 Sep 2026 11:48:36 +0200 Subject: [PATCH 1/3] Static-shape unconstrained rational-quadratic spline (torch.compile friendly) unconstrained_rational_quadratic_spline selected the inside/outside-tail points with boolean-mask indexing and a torch.any() branch. Those are data-dependent shapes and host synchronisations, which force graph breaks under torch.compile and keep the many small spline kernels from fusing; in a flow with hundreds of spline transforms this makes torch.compile a net slowdown. Evaluate the spline on the whole tensor with the inputs clamped into the domain and select the linear tails with torch.where instead. The result is bit-identical to the masked formulation (outputs, logabsdet and gradients). rational_quadratic_spline gains a check_domain flag (default True, unchanged behaviour) so the clamped call can skip the host-syncing domain check and discriminant assertion. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01JXDcmhZcz96xz9Zm475iAa --- .../transforms/splines/rational_quadratic.py | 91 +++++++----- .../splines/rational_quadratic_test.py | 137 ++++++++++++++++++ 2 files changed, 191 insertions(+), 37 deletions(-) diff --git a/nflows/transforms/splines/rational_quadratic.py b/nflows/transforms/splines/rational_quadratic.py index bfe57f3..c50ea09 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,51 @@ 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 + + # Clamping keeps every input inside the spline domain, so the (host-syncing) + # domain check of the inner spline is redundant here. + spline_outputs, spline_logabsdet = rational_quadratic_spline( + inputs=inputs.clamp(-tail_bound, tail_bound), + 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 +83,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 +150,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..ea90b27 100644 --- a/tests/transforms/splines/rational_quadratic_test.py +++ b/tests/transforms/splines/rational_quadratic_test.py @@ -82,3 +82,140 @@ 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.""" + import numpy as np + from torch.nn import functional as F + + from nflows.transforms.splines.rational_quadratic import ( + DEFAULT_MIN_DERIVATIVE, + rational_quadratic_spline, + ) + + 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 - 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 = 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 == "all_inside": + inputs = torch.rand(*shape, dtype=dtype) * 2 - 1 + elif mode == "boundary": + 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]): + # When every input lies in the tails the reference never touches the + # spline parameters (grad None); the static-shape version returns an + # explicit zero gradient, which is 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", "all_inside", "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): + if not hasattr(torch, "compile"): + self.skipTest("torch.compile not available") + inputs, params = self._make_inputs(torch.float32, "mixed") + for inverse in (False, True): + eager_out, eager_logabsdet = splines.unconstrained_rational_quadratic_spline( + inputs, *params, inverse=inverse + ) + compiled = torch.compile( + splines.unconstrained_rational_quadratic_spline, + fullgraph=True, + backend="aot_eager", + ) + out, logabsdet = compiled(inputs, *params, inverse=inverse) + self.eps = 1e-6 + self.assertEqual(out, eager_out) + self.assertEqual(logabsdet, eager_logabsdet) + + def test_inner_spline_domain_check_is_optional(self): + from nflows.transforms.base import InputOutsideDomain + + num_bins = 8 + shape = [4, 3] + params = [ + torch.randn(*shape, num_bins), + torch.randn(*shape, num_bins), + torch.randn(*shape, num_bins + 1), + ] + inputs = torch.rand(*shape) + 1.5 # outside [0, 1] + with self.assertRaises(InputOutsideDomain): + splines.rational_quadratic_spline(inputs, *params) + with self.assertRaises(InputOutsideDomain): + splines.rational_quadratic_spline(inputs, *params, check_domain=True) + # The caller takes responsibility for the domain; no host-side check. + splines.rational_quadratic_spline(inputs.clamp(0, 1), *params, check_domain=False) From 3257366f80ddcedd6a1c9b3747a169954d411323 Mon Sep 17 00:00:00 2001 From: Nihar Gupte Date: Fri, 11 Sep 2026 11:37:48 +0200 Subject: [PATCH 2/3] Trim the static-shape spline tests Keep the masked reference implementation as the oracle (outputs, log-det and gradients, mixed / all-outside / boundary inputs, float32 and float64, forward and inverse) and a fullgraph compile check; drop the check_domain unit test. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG --- .../splines/rational_quadratic_test.py | 69 +++++-------------- 1 file changed, 19 insertions(+), 50 deletions(-) diff --git a/tests/transforms/splines/rational_quadratic_test.py b/tests/transforms/splines/rational_quadratic_test.py index ea90b27..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): @@ -93,26 +96,18 @@ def _reference_unconstrained_rational_quadratic_spline( tail_bound=1.0, ): """The previous (masked-indexing) implementation, kept as a numerical reference.""" - import numpy as np - from torch.nn import functional as F - - from nflows.transforms.splines.rational_quadratic import ( - DEFAULT_MIN_DERIVATIVE, - rational_quadratic_spline, - ) - 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 - DEFAULT_MIN_DERIVATIVE) - 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 = rational_quadratic_spline( + 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, :], @@ -142,9 +137,7 @@ def _make_inputs(self, dtype, mode): elif mode == "all_outside": inputs = torch.sign(torch.randn(*shape)) * (1.0 + torch.rand(*shape)) inputs = inputs.to(dtype) - elif mode == "all_inside": - inputs = torch.rand(*shape, dtype=dtype) * 2 - 1 - elif mode == "boundary": + 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), @@ -169,9 +162,8 @@ def run(fn): 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]): - # When every input lies in the tails the reference never touches the - # spline parameters (grad None); the static-shape version returns an - # explicit zero gradient, which is equivalent. + # 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: @@ -179,43 +171,20 @@ def run(fn): def test_matches_reference_implementation(self): for dtype in (torch.float32, torch.float64): - for mode in ("mixed", "all_outside", "all_inside", "boundary"): + 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): - if not hasattr(torch, "compile"): - self.skipTest("torch.compile not available") inputs, params = self._make_inputs(torch.float32, "mixed") - for inverse in (False, True): - eager_out, eager_logabsdet = splines.unconstrained_rational_quadratic_spline( - inputs, *params, inverse=inverse - ) - compiled = torch.compile( - splines.unconstrained_rational_quadratic_spline, - fullgraph=True, - backend="aot_eager", - ) - out, logabsdet = compiled(inputs, *params, inverse=inverse) - self.eps = 1e-6 - self.assertEqual(out, eager_out) - self.assertEqual(logabsdet, eager_logabsdet) - - def test_inner_spline_domain_check_is_optional(self): - from nflows.transforms.base import InputOutsideDomain - - num_bins = 8 - shape = [4, 3] - params = [ - torch.randn(*shape, num_bins), - torch.randn(*shape, num_bins), - torch.randn(*shape, num_bins + 1), - ] - inputs = torch.rand(*shape) + 1.5 # outside [0, 1] - with self.assertRaises(InputOutsideDomain): - splines.rational_quadratic_spline(inputs, *params) - with self.assertRaises(InputOutsideDomain): - splines.rational_quadratic_spline(inputs, *params, check_domain=True) - # The caller takes responsibility for the domain; no host-side check. - splines.rational_quadratic_spline(inputs.clamp(0, 1), *params, check_domain=False) + 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]) From 20a62661e1d73b9ad2881a39729a0e3dbc35866e Mon Sep 17 00:00:00 2001 From: Nihar Gupte Date: Fri, 11 Sep 2026 12:10:31 +0200 Subject: [PATCH 3/3] Pass in-domain inputs through unchanged; clamp only the tails torch 2.14 changed the subgradient of clamp at its bounds from 1 to 0, so inputs lying exactly on the tail bound got a different input gradient than in the masked formulation (CI failure on the boundary parity test). Select the clamped value only for inputs outside the domain, as the masked version effectively does; values and graph shape are unchanged. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG --- nflows/transforms/splines/rational_quadratic.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/nflows/transforms/splines/rational_quadratic.py b/nflows/transforms/splines/rational_quadratic.py index c50ea09..352b3fe 100644 --- a/nflows/transforms/splines/rational_quadratic.py +++ b/nflows/transforms/splines/rational_quadratic.py @@ -42,10 +42,17 @@ def unconstrained_rational_quadratic_spline( unnormalized_derivatives[..., 0] = constant unnormalized_derivatives[..., -1] = constant - # Clamping keeps every input inside the spline domain, so the (host-syncing) - # domain check of the inner spline is redundant here. + # 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=inputs.clamp(-tail_bound, tail_bound), + inputs=spline_inputs, unnormalized_widths=unnormalized_widths, unnormalized_heights=unnormalized_heights, unnormalized_derivatives=unnormalized_derivatives,