Skip to content

Static-shape rational-quadratic spline (torch.compile friendly) - #13

Open
nihargupte-ph wants to merge 3 commits into
uofgravity:glasflowfrom
nihargupte-ph:compile-friendly-rqs
Open

nihargupte-ph wants to merge 3 commits into
uofgravity:glasflowfrom
nihargupte-ph:compile-friendly-rqs

Conversation

@nihargupte-ph

Copy link
Copy Markdown

unconstrained_rational_quadratic_spline selects the points inside and outside the tails with boolean-mask indexing (outputs[mask] = ...) and a torch.any branch. Both are data-dependent, so torch.compile breaks the graph in the coupling transform.

This PR rewrites the function with static shapes: the inputs are clamped into the domain, the spline is evaluated on the whole tensor, and the tails are selected with torch.where. The clamp happens before any spline arithmetic, so the unselected branch is finite and no inf or NaN leaks into the gradient. rational_quadratic_spline gets a check_domain=True kwarg so the caller can skip the host-syncing domain check when it has already clamped.

The result is numerically identical to the masked version. The test keeps the old implementation as a reference and checks outputs, log-det and gradients (mixed, all-outside and on-the-boundary inputs, float32 and float64, forward and inverse), plus a fullgraph=True compile check.

Motivation: in dingo this turns torch.compile on the NPE network from a slowdown into a 1.2 to 1.5x step-time gain, since the flow is bound by launching tens of thousands of small kernels per step.

Once this is merged, glasflow needs its submodules/nflows pointer bumped and a release.

nihargupte-ph and others added 2 commits September 2, 2026 11:48
…riendly)

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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01JXDcmhZcz96xz9Zm475iAa
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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG
@mj-will
mj-will self-requested a review September 11, 2026 10:04
@mj-will mj-will added the enhancement New feature or request label Sep 11, 2026
@mj-will

mj-will commented Sep 11, 2026

Copy link
Copy Markdown

Thanks @nihargupte-ph, I'll find some time to take a look over this. It looks like the CI is rather unhappy, but that may just be old Python version.

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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NCnvJRCVo7DBSU9YRiXcrG
@nihargupte-ph

Copy link
Copy Markdown
Author

Thanks @nihargupte-ph, I'll find some time to take a look over this. It looks like the CI is rather unhappy, but that may just be old Python version.

Right, I think locally my tests passed because I'm using torch=2.9. I think the difference is that the gradient is slightly different between old versions of pytorch and newer ones since in torch=2.14, clamp gives a gradient of 0 on the boundary. Just pushed a fix which will only apply the clamp to the tail

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants