Alternative #5 fix: NUFFT-level primitives + calibrated Pallas thresholds + cleanup - #7
Merged
Merged
Conversation
…port
The public NUFFT functions (nufft1d1...nufft3d3) were wrapped in @jax.custom_vjp,
whose primitive (custom_vjp_call) has no transpose rule. This blocked
jax.linear_transpose and therefore lax.custom_linear_solve, breaking
jax.scipy.sparse.linalg.{cg,gmres,bicgstab} on NUFFT-based operators (issue #5).
Replace with 9 jax.extend.core.Primitive instances, each carrying explicit JVP
and transpose rules. Type 1 and Type 2 transpose to each other (same isign);
Type 3 transposes to itself with source/target points swapped.
jax.grad now follows JAX-standard complex gradient convention (2*dL/dc-bar via
Wirtinger), which is the conjugate of the convention the previous custom_vjp
was returning. Tests updated accordingly.
The previous batching rule only handled vmap over the source argument (c/f), which broke the quickstart's vmap-over-x example. Add a generic fallback that vmap-traces the impl function for any other pattern (vmap over coordinates, multi-axis vmap, etc.). The fast path for source-only batching is preserved to keep zero overhead on the common case.
Adds spread_3d_pallas and interp_3d_pallas mirroring the existing 1D/2D implementations: triple loop over nspread offsets, separable kernel factorization, atomic scatter-add for spreading, gather+accumulate for interpolation. Wires _spread_3d_dispatch and _interp_3d_dispatch into the public spread_3d and interp_3d entry points, using the same M thresholds (10K for spread, 1M for interp). The backward VJP rules still call the pure-JAX impls directly — preserving existing behavior; tuning those paths is out of scope here. Tests are gpu_only-marked: collect on all platforms, skip on CPU, exercise correctness vs the pure-JAX reference plus the spread/interp adjoint inner-product identity on GPU.
The 3D spread/interp kernels used a Python-unrolled nspread^3 triple loop, which made Triton compile times explode (a tiny 256-point test never finished compiling on H100). Replace with three nested jax.lax.fori_loop levels: the IR stays small (one body per axis) and the z/y weights are hoisted to keep the separable factorization. A flat-index fori_loop was tried first but hit a Triton i32/i1 type assertion in the `_mod` lowering of the loop induction variable; nested loops avoid index decomposition entirely. Verified on H100: all 9 tests/test_pallas_3d.py pass, correctness rel_err ~1e-6 vs pure JAX. Speed: the kernels are currently 2-8x SLOWER than the pure-JAX path (interp especially, since XLA already fuses gather+FFT optimally), so the dispatch wiring should be revisited before relying on it — tracked separately. Also merges the 3D cases into benchmarks/bench_pallas_vs_jax.py (spread/interp/nufft3d1/nufft3d2) and removes the standalone benchmarks/bench_pallas_3d.py.
…ation
Measured cleanly on H100 (the earlier "Pallas is slow" conclusion was a
benchmark artifact: a poisoned _HAS_PALLAS_GPU toggle + cached primitive
lowering). Reality:
spread pure-JAX impl vs Pallas (M=1M): 1D 56x 2D 20x 3D 4x faster
end-to-end public NUFFT (default path, H100):
Type1 1d1(256)=1.43ms 2d1(256^2)=2.88ms @ M=1M
1d1=14.6ms 2d1=27ms @ M=10M
Type2 1d2=0.10ms 2d2=0.30ms 3d2=3.0ms @ M=1M
That is FINUFFT/cufinufft-class; the existing Triton spread kernels were
already the win. So:
- Keep spread_{1,2,3}d_pallas + their dispatch (3-56x, FINUFFT-class).
- Remove the interp Pallas kernels and route all interp to pure XLA: it
was at best parity (1D/2D) and 2-4x SLOWER in 3D, so dropping it is a
real Type-2/3D speedup and less code. _PALLAS_MIN_M_INTERP removed.
- Delete the dead exploratory subproblem reformulation (pure-JAX
binned/tiled spread: proven non-competitive vs the Triton kernels) and
consolidate benchmarks into one streaming bench_pallas_vs_jax.py.
- test_pallas_3d.py trimmed to the surviving spread kernel.
CPU suite: 215 passed, 4 skipped (GPU-only spread-3d). 3D Pallas spread
verified passing on H100 earlier in this work.
Measured on H100 (clean streaming bench): the original
_PALLAS_MIN_M_SPREAD=10_000 was a single value missing the per-dim crossover.
1D crossover ~ M=5_000 (impl 0.85x at M=2K -> Pallas 1.5x at M=5K -> 33x at M=100K)
2D crossover ~ M=5_000 (impl 0.50x at M=2K -> Pallas 1.1x at M=5K -> 31x at M=100K)
3D crossover ~ M=10_000+ (impl wins below; Pallas 5.6x at M=50K)
Split into _PALLAS_MIN_M_SPREAD_{1D,2D,3D} = 5_000, 5_000, 10_000 so 1D/2D
catch the M=5K..10K win (~1.5-2.6x) without activating Pallas in the regime
where kernel-launch and BLOCK_SIZE padding still dominate (small M).
Sweeps M in {100..100K} for spread_{1d,2d,3d}, impl vs Pallas, on whatever
device is present. Used to derive _PALLAS_MIN_M_SPREAD_{1D,2D,3D} = 5K, 5K, 10K
on H100. Re-run after a GPU change to recalibrate.
After 'git merge origin/main -X ours' some non-conflicting hunks from PR #6 landed and clashed with our design. Clean up: - pallas_spread.py: drop PR #6's dead imports (Primitive, ShapedArray, ad, batching, _PALLAS_MIN_M_INTERP etc) — none used since we keep our primitives at the NUFFT level, not the spread level - spread.py dispatch: spread_1d/2d_primitive -> spread_1d/2d_pallas (PR #6 routed via primitives that don't exist on our branch) - transforms/nufft1/2/3.py: spread_3d_impl -> _spread_3d_dispatch (and interp_3d_impl -> _interp_3d_dispatch) so the 3D Pallas dispatch our branch added is not silently bypassed. CPU tests: 215 passed, 4 skipped.
Contributor
|
Great job! Right now the logic of autodiff is much more clearer. |
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.
Highlights
nufft{1,2,3}d{1,2,3}areprimitives, with Type 1 ↔ Type 2 transpose same
isign, Type 3self-adjoint with source/target swapped.
autodiff.pyshrinks ~1500 lines.runs): the single
_PALLAS_MIN_M_SPREAD=10_000was miscalibrated. New{1D: 5_000, 2D: 5_000, 3D: 10_000}— 1D/2D crossover is M=5K, not 10K.parity in 1D/2D, 2–4× slower in 3D — removing them is a real Type-2/3D
speedup and less code.
fori_loopto avoid thenspread³triple-unroll compile blowup).jax.jvp(nufft*)now works directly (was blocked bycustom_vjp).