Skip to content

Alternative #5 fix: NUFFT-level primitives + calibrated Pallas thresholds + cleanup - #7

Merged
geoffroyO merged 12 commits into
mainfrom
feat/jax-primitives
Jun 5, 2026
Merged

geoffroyO merged 12 commits into
mainfrom
feat/jax-primitives

Conversation

@geoffroyO

@geoffroyO geoffroyO commented May 20, 2026 •

Copy link
Copy Markdown
Collaborator

Highlights

  • Primitives at the public NUFFT level: nufft{1,2,3}d{1,2,3} are
    primitives, with Type 1 ↔ Type 2 transpose same isign, Type 3
    self-adjoint with source/target swapped. autodiff.py shrinks ~1500 lines.
  • Per-dim Pallas spread thresholds calibrated on H100 (median of 30 warm
    runs): the single _PALLAS_MIN_M_SPREAD=10_000 was miscalibrated. New
    {1D: 5_000, 2D: 5_000, 3D: 10_000} — 1D/2D crossover is M=5K, not 10K.
  • Drop the Pallas interp kernels and route interp to pure XLA: measured
    parity in 1D/2D, 2–4× slower in 3D — removing them is a real Type-2/3D
    speedup and less code.
  • 3D Pallas spread + tests + dispatch added (kernel uses nested Triton
    fori_loop to avoid the nspread³ triple-unroll compile blowup).
  • Public jax.jvp(nufft*) now works directly (was blocked by custom_vjp).

geoffroyO added 11 commits May 13, 2026 14:05
…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.
@gRox167

gRox167 commented May 20, 2026

Copy link
Copy Markdown
Contributor

Great job! Right now the logic of autodiff is much more clearer.

@geoffroyO
geoffroyO merged commit fc09095 into main Jun 5, 2026
3 checks passed
@geoffroyO
geoffroyO deleted the feat/jax-primitives branch June 5, 2026 09:40
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants