Static-shape rational-quadratic spline (torch.compile friendly) - #13
Open
nihargupte-ph wants to merge 3 commits into
Open
nihargupte-ph wants to merge 3 commits into
nihargupte-ph wants to merge 3 commits into
Conversation
…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
self-requested a review
September 11, 2026 10:04
|
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
Author
Right, I think locally my tests passed because I'm using |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
unconstrained_rational_quadratic_splineselects the points inside and outside the tails with boolean-mask indexing (outputs[mask] = ...) and atorch.anybranch. Both are data-dependent, sotorch.compilebreaks 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_splinegets acheck_domain=Truekwarg 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.