jit compile tangent computation in Proximal - #2239
Conversation
Memory benchmark result| Test Name | %Δ | Master (MB) | PR (MB) | Δ (MB) | Time PR (s) | Time Master (s) |
| -------------------------------------- | ------------ | ------------------ | ------------------ | ------------ | ------------------ | ------------------ |
test_objective_jac_w7x | -0.16 % | 4.227e+03 | 4.220e+03 | -6.81 | 31.40 | 29.14 |
test_proximal_jac_w7x_with_eq_update | -4.51 % | 6.854e+03 | 6.545e+03 | -309.19 | 148.84 | 154.65 |
test_proximal_freeb_jac | -0.74 % | 1.354e+04 | 1.344e+04 | -100.71 | 76.85 | 80.79 |
test_proximal_freeb_jac_blocked | -1.31 % | 7.887e+03 | 7.784e+03 | -103.34 | 64.69 | 72.36 |
test_proximal_freeb_jac_batched | -0.83 % | 7.859e+03 | 7.794e+03 | -65.13 | 64.00 | 72.75 |
test_proximal_jac_ripple | -0.94 % | 3.791e+03 | 3.755e+03 | -35.49 | 49.96 | 57.38 |
test_proximal_jac_ripple_bounce1d | -1.13 % | 3.974e+03 | 3.929e+03 | -44.79 | 63.95 | 72.34 |
test_eq_solve | 1.10 % | 1.832e+03 | 1.852e+03 | 20.23 | 52.23 | 54.86 |
test_objective_quadratic_flux_jac | -0.02 % | 1.888e+03 | 1.888e+03 | -0.38 | 34.34 | 37.15 |For the memory plots, go to the summary of |
| benchmark_name | dt(%) | dt(s) | t_new(s) | t_old(s) |
| -------------------------------------- | ---------------------- | ---------------------- | ---------------------- | ---------------------- |
test_build_transform_fft_lowres | -3.96 +/- 3.14 | -3.11e-02 +/- 2.47e-02 | 7.54e-01 +/- 2.0e-02 | 7.85e-01 +/- 1.4e-02 |
test_equilibrium_init_lowres | -2.79 +/- 4.14 | -1.68e-01 +/- 2.49e-01 | 5.85e+00 +/- 1.6e-01 | 6.02e+00 +/- 1.9e-01 |
test_objective_compile_atf | -0.88 +/- 2.30 | -4.58e-02 +/- 1.20e-01 | 5.16e+00 +/- 5.8e-02 | 5.21e+00 +/- 1.0e-01 |
test_objective_compute_atf | +0.59 +/- 8.71 | +1.22e-05 +/- 1.81e-04 | 2.09e-03 +/- 1.2e-04 | 2.08e-03 +/- 1.4e-04 |
test_objective_jac_atf | +0.04 +/- 2.17 | +6.54e-04 +/- 3.24e-02 | 1.49e+00 +/- 2.4e-02 | 1.49e+00 +/- 2.2e-02 |
test_perturb_1 | -1.54 +/- 3.41 | -1.64e-01 +/- 3.63e-01 | 1.05e+01 +/- 2.5e-01 | 1.06e+01 +/- 2.6e-01 |
test_proximal_jac_atf | +2.27 +/- 2.77 | +1.05e-01 +/- 1.29e-01 | 4.75e+00 +/- 9.8e-02 | 4.65e+00 +/- 8.3e-02 |
test_proximal_jac_atf_with_eq_update | +0.70 +/- 3.03 | +7.31e-02 +/- 3.18e-01 | 1.06e+01 +/- 1.3e-01 | 1.05e+01 +/- 2.9e-01 |
test_proximal_freeb_compute | +2.45 +/- 8.17 | +2.99e-03 +/- 9.95e-03 | 1.25e-01 +/- 8.2e-03 | 1.22e-01 +/- 5.6e-03 |
test_solve_fixed_iter_compiled | +0.76 +/- 2.81 | +4.45e-02 +/- 1.64e-01 | 5.89e+00 +/- 1.5e-01 | 5.84e+00 +/- 7.3e-02 |
test_LinearConstraintProjection_build | -1.40 +/- 2.40 | -8.04e-02 +/- 1.37e-01 | 5.65e+00 +/- 6.5e-02 | 5.73e+00 +/- 1.2e-01 |
test_objective_compute_ripple | +0.95 +/- 5.08 | +1.79e-03 +/- 9.52e-03 | 1.89e-01 +/- 6.1e-03 | 1.87e-01 +/- 7.3e-03 |
test_objective_grad_ripple | +2.61 +/- 3.46 | +2.29e-02 +/- 3.04e-02 | 9.00e-01 +/- 1.9e-02 | 8.77e-01 +/- 2.4e-02 |
test_objective_quadratic_flux_compute | -3.62 +/- 3.36 | -9.67e-04 +/- 8.97e-04 | 2.57e-02 +/- 7.0e-04 | 2.67e-02 +/- 5.5e-04 |
test_build_transform_fft_midres | -1.08 +/- 4.59 | -9.59e-03 +/- 4.10e-02 | 8.82e-01 +/- 3.6e-02 | 8.92e-01 +/- 2.0e-02 |
test_build_transform_fft_highres | -1.28 +/- 3.37 | -1.52e-02 +/- 3.99e-02 | 1.17e+00 +/- 2.5e-02 | 1.18e+00 +/- 3.1e-02 |
test_equilibrium_init_medres | -0.27 +/- 2.17 | -1.88e-02 +/- 1.50e-01 | 6.87e+00 +/- 1.2e-01 | 6.89e+00 +/- 9.3e-02 |
test_objective_compile_dshape_current | +0.97 +/- 2.82 | +4.04e-02 +/- 1.17e-01 | 4.20e+00 +/- 9.8e-02 | 4.16e+00 +/- 6.5e-02 |
test_objective_compute_dshape_current | +5.48 +/- 10.75 | +3.86e-05 +/- 7.57e-05 | 7.42e-04 +/- 5.9e-05 | 7.04e-04 +/- 4.8e-05 |
test_objective_jac_dshape_current | -5.96 +/- 26.95 | -1.54e-03 +/- 6.97e-03 | 2.43e-02 +/- 5.7e-03 | 2.59e-02 +/- 4.0e-03 |
test_perturb_2 | -1.81 +/- 1.93 | -2.86e-01 +/- 3.06e-01 | 1.56e+01 +/- 1.2e-01 | 1.58e+01 +/- 2.8e-01 |
+test_proximal_jac_atf_chunked | -67.92 +/- 0.68 | -1.14e+01 +/- 1.14e-01 | 5.38e+00 +/- 3.7e-02 | 1.68e+01 +/- 1.1e-01 |
test_proximal_freeb_jac | +0.72 +/- 2.85 | +3.44e-02 +/- 1.37e-01 | 4.83e+00 +/- 1.2e-01 | 4.80e+00 +/- 7.1e-02 |
test_solve_fixed_iter | -0.64 +/- 2.43 | -1.56e-01 +/- 5.90e-01 | 2.42e+01 +/- 4.5e-01 | 2.43e+01 +/- 3.8e-01 |
test_objective_compute_ripple_bounce1d | -1.63 +/- 3.96 | -5.01e-03 +/- 1.22e-02 | 3.02e-01 +/- 9.6e-03 | 3.07e-01 +/- 7.5e-03 |
test_objective_grad_ripple_bounce1d | +1.17 +/- 2.32 | +1.10e-02 +/- 2.20e-02 | 9.58e-01 +/- 2.1e-02 | 9.47e-01 +/- 6.4e-03 |
test_objective_quadratic_flux_jac | +0.81 +/- 1.23 | +1.80e-02 +/- 2.73e-02 | 2.25e+00 +/- 1.6e-02 | 2.23e+00 +/- 2.2e-02 |Github CI performance can be noisy. When evaluating the benchmarks, developers should take this into account. |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #2239 +/- ##
==========================================
- Coverage 94.35% 94.35% -0.01%
==========================================
Files 101 101
Lines 29036 29034 -2
==========================================
- Hits 27396 27394 -2
Misses 1640 1640
🚀 New features to boost your workflow:
|
…t fused into a giant program
| # Note: Fxh and its SVD do not depend on dc (the vectorized argument). Since the | ||
| # whole tangent computation is jitted as one program, we rely on the compiler to | ||
| # hoist this loop-invariant SVD out of the batched scan/vmap rather than | ||
| # recomputing it for every tangent. |
There was a problem hiding this comment.
Previously, I was trying to be smart by removing the transpose from the Fxh and handle that at the return operation (transpose of u and v etc) but that cause slowdown on CPU. Maybe SVD lowering of JAX on CPU is dependent on tall/wide?!
Bumps [actions/setup-python](https://github.com/actions/setup-python) from 6 to 7. <details> <summary>Release notes</summary> <p><em>Sourced from <a href="https://github.com/actions/setup-python/releases">actions/setup-python's releases</a>.</em></p> <blockquote> <h2>v7.0.0</h2> <h2>What's Changed</h2> <h3>Enhancements</h3> <ul> <li>Migrate to ESM and upgrade dependencies by <a href="https://github.com/priyagupta108"><code>@priyagupta108</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1330">actions/setup-python#1330</a></li> <li>Pin SHA commits and update docs with latest versions by <a href="https://github.com/HarithaVattikuti"><code>@HarithaVattikuti</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1338">actions/setup-python#1338</a></li> <li>Remove the pip-install input by <a href="https://github.com/gowridurgad"><code>@gowridurgad</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1336">actions/setup-python#1336</a></li> </ul> <h3>Bug Fix</h3> <ul> <li>Fix to Classify stderr warning messages as warnings instead of errors in annotations by <a href="https://github.com/lmvysakh"><code>@lmvysakh</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1335">actions/setup-python#1335</a></li> <li>Validate and retry manifest fetch to prevent silent failures by <a href="https://github.com/priyagupta108"><code>@priyagupta108</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1332">actions/setup-python#1332</a></li> </ul> <h3>Dependency Upgrade</h3> <ul> <li>Bump certifi from 2020.6.20 to 2024.7.4 in /<strong>tests</strong>/data by <a href="https://github.com/dependabot"><code>@dependabot</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1328">actions/setup-python#1328</a></li> <li>Remove EOL Python versions and Bumps numpy text fixture by <a href="https://github.com/priya-kinthali"><code>@priya-kinthali</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1333">actions/setup-python#1333</a></li> <li>Upgrade <code>@actions/cache</code> to 6.2.0 by <a href="https://github.com/philip-gai"><code>@philip-gai</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1337">actions/setup-python#1337</a></li> </ul> <h2>New Contributors</h2> <ul> <li><a href="https://github.com/lmvysakh"><code>@lmvysakh</code></a> made their first contribution in <a href="https://redirect.github.com/actions/setup-python/pull/1335">actions/setup-python#1335</a></li> <li><a href="https://github.com/philip-gai"><code>@philip-gai</code></a> made their first contribution in <a href="https://redirect.github.com/actions/setup-python/pull/1337">actions/setup-python#1337</a></li> </ul> <p><strong>Full Changelog</strong>: <a href="https://github.com/actions/setup-python/compare/v6...v7.0.0">https://github.com/actions/setup-python/compare/v6...v7.0.0</a></p> <h2>v6.3.0</h2> <h2>What's Changed</h2> <h3>Enhancement</h3> <ul> <li>Add RHEL support and include Linux distro in cache keys by <a href="https://github.com/priyagupta108"><code>@priyagupta108</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1323">actions/setup-python#1323</a></li> <li>Fix pip cache error handling on Windows by <a href="https://github.com/priyagupta108"><code>@priyagupta108</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1040">actions/setup-python#1040</a></li> </ul> <h3>Dependency update</h3> <ul> <li>Upgrade minimatch from 3.1.2 to 3.1.5 by <a href="https://github.com/dependabot"><code>@dependabot</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1281">actions/setup-python#1281</a></li> <li>Upgrade actions dependencies by <a href="https://github.com/gowridurgad"><code>@gowridurgad</code></a> with <a href="https://github.com/Copilot"><code>@Copilot</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1303">actions/setup-python#1303</a></li> <li>Upgrade <code>@actions/cache</code> to 5.1.0, log cache write denied by <a href="https://github.com/jasongin"><code>@jasongin</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1324">actions/setup-python#1324</a></li> <li>Upgrade dependency versions and test workflow configuration by <a href="https://github.com/HarithaVattikuti"><code>@HarithaVattikuti</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1322">actions/setup-python#1322</a></li> </ul> <h3>Documentation</h3> <ul> <li>Update advanced-usage.md by <a href="https://github.com/Dunky-Z"><code>@Dunky-Z</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/811">actions/setup-python#811</a></li> </ul> <h2>New Contributors</h2> <ul> <li><a href="https://github.com/gowridurgad"><code>@gowridurgad</code></a> with <a href="https://github.com/Copilot"><code>@Copilot</code></a> made their first contribution in <a href="https://redirect.github.com/actions/setup-python/pull/1303">actions/setup-python#1303</a></li> <li><a href="https://github.com/jasongin"><code>@jasongin</code></a> made their first contribution in <a href="https://redirect.github.com/actions/setup-python/pull/1324">actions/setup-python#1324</a></li> <li><a href="https://github.com/Dunky-Z"><code>@Dunky-Z</code></a> made their first contribution in <a href="https://redirect.github.com/actions/setup-python/pull/811">actions/setup-python#811</a></li> </ul> <p><strong>Full Changelog</strong>: <a href="https://github.com/actions/setup-python/compare/v6.2.0...v6.3.0">https://github.com/actions/setup-python/compare/v6.2.0...v6.3.0</a></p> <h2>v6.2.0</h2> <h2>What's Changed</h2> <h3>Dependency Upgrades</h3> <ul> <li>Upgrade dependencies to Node 24 compatible versions by <a href="https://github.com/salmanmkc"><code>@salmanmkc</code></a> in <a href="https://redirect.github.com/actions/setup-python/pull/1259">actions/setup-python#1259</a></li> </ul> <!-- raw HTML omitted --> </blockquote> <p>... (truncated)</p> </details> <details> <summary>Commits</summary> <ul> <li><a href="https://github.com/actions/setup-python/commit/5fda3b95a4ea91299a34e894583c3862153e4b97"><code>5fda3b9</code></a> Pin SHA commits and update docs with latest versions (<a href="https://redirect.github.com/actions/setup-python/issues/1338">#1338</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/4ab7e95f05e168b4356aebde89dd84f59c283d8e"><code>4ab7e95</code></a> Merge pull request <a href="https://redirect.github.com/actions/setup-python/issues/1337">#1337</a> from actions/philip-gai/bump-actions-cache-6-2-0</li> <li><a href="https://github.com/actions/setup-python/commit/0f3a009f475dbea83c0371cd85d099690fee8c5c"><code>0f3a009</code></a> Remove the pip-install input (<a href="https://redirect.github.com/actions/setup-python/issues/1336">#1336</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/f8cf4291c8b8e273ddd26e569454615c7315d932"><code>f8cf429</code></a> Migrate to ESM and upgrade dependencies (<a href="https://redirect.github.com/actions/setup-python/issues/1330">#1330</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/54baeea5b34417d10a7479663a23cca53ea209b5"><code>54baeea</code></a> Validate and retry manifest fetch to prevent silent failures (<a href="https://redirect.github.com/actions/setup-python/issues/1332">#1332</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/c7092773a316760f4ecfe498e4af668a4dafeac5"><code>c709277</code></a> Annotation code fix (<a href="https://redirect.github.com/actions/setup-python/issues/1335">#1335</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/6849080452e69b330395e8a6d23cf90f56d76a1a"><code>6849080</code></a> remove EOL Python versions and Bumps numpy text fixture (<a href="https://redirect.github.com/actions/setup-python/issues/1333">#1333</a>)</li> <li><a href="https://github.com/actions/setup-python/commit/0903b469fbf4441aadfe4f4b249dc5b1fba3a73e"><code>0903b46</code></a> Bump certifi from 2020.6.20 to 2024.7.4 in /<strong>tests</strong>/data (<a href="https://redirect.github.com/actions/setup-python/issues/1328">#1328</a>)</li> <li>See full diff in <a href="https://github.com/actions/setup-python/compare/v6...v7">compare view</a></li> </ul> </details> <br /> [](https://docs.github.com/en/github/managing-security-vulnerabilities/about-dependabot-security-updates#about-compatibility-scores) Dependabot will resolve any conflicts with this PR as long as you don't alter it yourself. You can also trigger a rebase manually by commenting `@dependabot rebase`. [//]: # (dependabot-automerge-start) [//]: # (dependabot-automerge-end) --- <details> <summary>Dependabot commands and options</summary> <br /> You can trigger Dependabot actions by commenting on this PR: - `@dependabot rebase` will rebase this PR - `@dependabot recreate` will recreate this PR, overwriting any edits that have been made to it - `@dependabot show <dependency name> ignore conditions` will show all of the ignore conditions of the specified dependency - `@dependabot ignore this major version` will close this PR and stop Dependabot creating any more for this major version (unless you reopen the PR or upgrade to it yourself) - `@dependabot ignore this minor version` will close this PR and stop Dependabot creating any more for this minor version (unless you reopen the PR or upgrade to it yourself) - `@dependabot ignore this dependency` will close this PR and stop Dependabot creating any more for this dependency (unless you reopen the PR or upgrade to it yourself) </details> Signed-off-by: dependabot[bot] <support@github.com> Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
…w by rescale (#2268) In our reactor QA script, we use `FixCurrent` assuming the [0,1] indices are the constant and rho^1 modes of the profile. However after the changes in #1871, rescaling an eq now returns `ScaledProfile` object which means now that [0,1] indices are the scale and the constant mode of the power series. this resulted in letting the linear term be free, and a nonzero derivative of current at the axis which is unphysical (this also leads to unphysical iota profiles due to the way the axis limit of iota works compared to how iota is computed away from rho=0). This fixes that unintentional bug and update the reactor_QA output accordingly. The reactor_QA example .h5 file had this issue since #1907 updated it (my bad...). <img width="276" height="276" alt="image" src="https://github.com/user-attachments/assets/83d429f3-f4e5-4ade-a700-ade8cc2ae96e" /> --------- Co-authored-by: daniel-dudt <daniel.dudt@princetonstellarators.energy> Co-authored-by: Yigit Gunsur Elmacioglu <102380275+YigitElma@users.noreply.github.com> Co-authored-by: YigitElma <yigitelmacioglu@gmail.com>
Notes that interpax>=0.3.14 Resolves #2120 and adds test anyways for it --------- Co-authored-by: Yigit Gunsur Elmacioglu <102380275+YigitElma@users.noreply.github.com>
…nking fix PR PlasmaControl#2239 (upstream) restructures ProximalProjection's tangent computation to form the reduced constraint Jacobian densely once, which is a strictly better fix for the SVD-per-chunk problem than linearizing the per-direction tangent function. Revert _constraint_wrappers.py to its pre-branch state so there's no overlap/conflict with that PR, and remove the two tests that exercised the now-dropped consolidation. The general Derivative.linearize fix in objective_funs.py/batching.py is unrelated and still valuable on its own (and is load-bearing for PlasmaControl#2239's own single jvp_op call once chunked), so it stays. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01JgWTyAewaigF78JQ6gH6Ry
| vs = jnp.split(v, np.cumsum(dimc_per_thing)[:-1], axis=-1) | ||
| # JVPs are taken for the dxdc @ v tangents, but at most dimc of them are useful, | ||
| # so take whichever combination needs fewer of them. Applying v after the JVPs | ||
| # only pays off with a lot of coil, surface etc DoFs, ie. single stage. |
There was a problem hiding this comment.
I don't think I understand the difference bwtn these two paths
There was a problem hiding this comment.
dim_opt= number of free optimization variables (Rb+Zb+profiles+coils etc after linear constraints)prox.dim_x= total number of optimization variables (Rb+Zb+profiles+coils etc)dim_c= total number of eq optimization variables (Rb+Zb+profiles)dim_xeq= size of full eq state vector (Rlmn+Zlmn+Llmn+Rb+Zb etc)
_proximal_eq_tangents calls the JVP of the constraint with eq_feasible_tangents and tangents for the optimization variables. vs is a dim_opt by prox.dim_x sized array, when we take vs[eq_idx], it becomes dim_opt by dim_c. So, the indexing operation doesn't change the number of tangents but makes the tangent vectors shorter.
We know that JVP wrt non-equilibrium optimization variables are 0, so if dim_opt > dim_c, we are going to have some trivially 0 JVPs. We can either take JVP wrt to vs[eq_idx] @ dxdc.T which is dim_opt by dim_xeq or wrt dxdc.T which is dim_c by dim_xeq. These branches compute the lower one with a static check.
There was a problem hiding this comment.
I originilly had a comment like this
# dxdc maps the eq DoFs (Rb_lmn, Zb_lmn, profiles) to the full eq state vector.
# Everything in _proximal_eq_tangents is linear in the directions it gets, so we
# can either take the JVPs for the columns of dxdc and apply v afterwards with a
# matmul (first branch), or apply v first and take the JVPs only for the
# directions in v (second branch). The JVPs dominate the cost, so we use the one
# with fewer directions. The first branch never needs more than
# dimc_per_thing[eq_idx] of them, since that is all the eq DoFs there are, so it
# only pays off when v has more directions than that. Usually v holds the feasible
# tangents of the outer LinearConstraintProjection and we optimize only some of
# the boundary modes, so the second branch is the usual one and the first only
# kicks in with a lot of coil, surface etc DoFs, important for single stage.
But it was very long, I can keep it if this is not clear.
batched_vectorizecall from Proximal, instead conducting the vectorized operation manuallyForceBalancefor each tangent, only computes for a maximum ofdimc_per_thing[eq_idx]tangents. For example, lets say you have 300 boundary coefficients and 100 coil coefficients inside a free boundary solve, master computes 400 JVPs, this PR computes 300.Note: Removing the
batched_vectorizewill help to make #1495 work with parallel force balance constraints too.The document I shared in #1623 (comment) might help to review the changes in this PR. One change from that document is that in this PR, I removed the
prox._feasible_tangentsand instead multiplied the related part byeq_feasible_tangents, these are equivalent because D@Z is exactlyeq_feasible_tangents.