Sparse pullback for big performance gain - #2170
Conversation
Co-authored-by: Kaya Unalmis <kayaunalmis@proton.me>
Co-authored-by: Kaya Unalmis <kayaunalmis@proton.me>
|
@dpanici when will this be merged? |
|
@ddudt when will this be merged |
|
It's needed for #2257 |
|
@f0uriest when will this be merged? |
| benchmark_name | dt(%) | dt(s) | t_new(s) | t_old(s) |
| -------------------------------------- | ---------------------- | ---------------------- | ---------------------- | ---------------------- |
test_build_transform_fft_lowres | +0.73 +/- 5.28 | +6.57e-03 +/- 4.76e-02 | 9.08e-01 +/- 2.7e-02 | 9.01e-01 +/- 3.9e-02 |
test_equilibrium_init_lowres | +1.81 +/- 3.35 | +1.25e-01 +/- 2.31e-01 | 7.01e+00 +/- 1.4e-01 | 6.88e+00 +/- 1.8e-01 |
test_objective_compile_atf | +2.59 +/- 4.30 | +1.60e-01 +/- 2.66e-01 | 6.35e+00 +/- 2.0e-01 | 6.19e+00 +/- 1.8e-01 |
test_objective_compute_atf | -4.19 +/- 8.26 | -9.55e-05 +/- 1.88e-04 | 2.19e-03 +/- 1.6e-04 | 2.28e-03 +/- 9.9e-05 |
test_objective_jac_atf | +2.76 +/- 3.12 | +4.38e-02 +/- 4.95e-02 | 1.63e+00 +/- 4.5e-02 | 1.59e+00 +/- 2.0e-02 |
test_perturb_1 | +3.04 +/- 1.51 | +3.76e-01 +/- 1.86e-01 | 1.27e+01 +/- 5.9e-02 | 1.23e+01 +/- 1.8e-01 |
test_proximal_jac_atf | +0.87 +/- 2.23 | +4.62e-02 +/- 1.18e-01 | 5.34e+00 +/- 7.1e-02 | 5.29e+00 +/- 9.5e-02 |
test_proximal_freeb_compute | +0.85 +/- 3.20 | +1.45e-03 +/- 5.44e-03 | 1.71e-01 +/- 4.1e-03 | 1.70e-01 +/- 3.6e-03 |
test_solve_fixed_iter | +1.86 +/- 3.22 | +4.68e-01 +/- 8.12e-01 | 2.57e+01 +/- 7.9e-01 | 2.52e+01 +/- 2.1e-01 |
test_LinearConstraintProjection_build | +2.53 +/- 7.46 | +1.82e-01 +/- 5.37e-01 | 7.38e+00 +/- 2.6e-01 | 7.20e+00 +/- 4.7e-01 |
test_objective_compute_ripple | +8.71 +/- 5.87 | +1.95e-02 +/- 1.31e-02 | 2.43e-01 +/- 6.2e-03 | 2.24e-01 +/- 1.2e-02 |
test_objective_quadratic_flux_compute | +1.04 +/- 12.99 | +5.40e-04 +/- 6.72e-03 | 5.23e-02 +/- 4.4e-03 | 5.18e-02 +/- 5.1e-03 |
test_build_transform_fft_midres | +0.26 +/- 4.57 | +2.51e-03 +/- 4.33e-02 | 9.50e-01 +/- 3.3e-02 | 9.47e-01 +/- 2.8e-02 |
test_build_transform_fft_highres | +0.40 +/- 2.81 | +4.93e-03 +/- 3.47e-02 | 1.24e+00 +/- 2.3e-02 | 1.24e+00 +/- 2.6e-02 |
test_equilibrium_init_medres | -0.90 +/- 3.93 | -6.61e-02 +/- 2.88e-01 | 7.27e+00 +/- 2.5e-01 | 7.33e+00 +/- 1.4e-01 |
test_objective_compile_dshape_current | -1.98 +/- 4.16 | -8.73e-02 +/- 1.83e-01 | 4.31e+00 +/- 1.3e-01 | 4.40e+00 +/- 1.3e-01 |
test_objective_compute_dshape_current | +5.30 +/- 13.60 | +4.01e-05 +/- 1.03e-04 | 7.97e-04 +/- 8.9e-05 | 7.57e-04 +/- 5.2e-05 |
test_objective_jac_dshape_current | +4.67 +/- 22.88 | +1.18e-03 +/- 5.79e-03 | 2.65e-02 +/- 4.5e-03 | 2.53e-02 +/- 3.6e-03 |
test_perturb_2 | -1.30 +/- 1.93 | -2.19e-01 +/- 3.24e-01 | 1.66e+01 +/- 1.9e-01 | 1.68e+01 +/- 2.6e-01 |
test_proximal_jac_atf_with_eq_update | +1.41 +/- 3.03 | +1.76e-01 +/- 3.79e-01 | 1.27e+01 +/- 3.6e-01 | 1.25e+01 +/- 1.1e-01 |
test_proximal_freeb_jac | +3.59 +/- 4.19 | +1.74e-01 +/- 2.03e-01 | 5.03e+00 +/- 1.3e-01 | 4.85e+00 +/- 1.6e-01 |
test_solve_fixed_iter_compiled | +0.95 +/- 4.41 | +5.86e-02 +/- 2.71e-01 | 6.19e+00 +/- 2.5e-01 | 6.13e+00 +/- 1.0e-01 |
test_objective_grad_ripple | +0.32 +/- 3.12 | +2.83e-03 +/- 2.79e-02 | 8.95e-01 +/- 2.4e-02 | 8.93e-01 +/- 1.4e-02 |
test_objective_quadratic_flux_jac | -0.06 +/- 1.13 | -5.44e-03 +/- 9.66e-02 | 8.56e+00 +/- 4.2e-02 | 8.57e+00 +/- 8.7e-02 |Github CI performance can be noisy. When evaluating the benchmarks, developers should take this into account. |
dpanici
left a comment
There was a problem hiding this comment.
Thanks for all the work that went into this PR! I would like these cosmetic changes made. I can do them myself if you prefer.
Main blocking comment I have is on the removal of jac_chunk_size: I get that these changes make the code much faster than before and make jac_chunk_size obsolete. It is good practice to then just throw a DeprecationWarning that it is no longer used or needed for this objective given these changes. I can make that commit as well (already made the suggestion in a comment).
The others are spelling or small warnings.
again, if you prefer I can make these changes myself as they are cosmetic. These are my main comments.
The other thought is that the Options class is a great idea and maybe could be used elsewhere outside of Bounce in the future, but I will leave that for a future problem.
| @@ -723,10 +809,10 @@ def theta_on_fieldlines(angle, iota, alpha, num_transit, NFP, *, X_min=24): | |||
| alpha : jnp.ndarray | |||
There was a problem hiding this comment.
does this function or elsewhere in Bounce2D require that the alphas be the same across all rho surfaces?
Just wondering, if not, we could allow them to be different per surface and that would aid #2257 . Not suggesting we do this in this PR or anything like that, just asking out of informational reasons for future reference.
| DeprecationWarning, | ||
| ) | ||
| bad_kwargs = kwargs.keys() - allowed_kwargs | ||
| bad_kwargs = kwargs.keys() - allowed_kwargs - {"num_transit"} |
There was a problem hiding this comment.
Should we just add this to allowed_kwargs? Or maybe make a new set called deprecated_kwargs to keep track of stuff like this?
There was a problem hiding this comment.
Im fine with anything
| normalize_target=True, | ||
| loss_function=None, | ||
| deriv_mode="auto", | ||
| jac_chunk_size=None, |
There was a problem hiding this comment.
Now that we are supporting forward mode I think we need to keep the jac_chunk_size argument.
| surf_batch_size=1, | ||
| nufft_eps=1e-7, | ||
| spline=True, | ||
| use_bounce1d=False, |
| return grid.meshgrid_reshape(f, "raz") | ||
|
|
||
| def points(self, pitch_inv, num_well=None): | ||
| def points(self, pitch_inv, num_well=-1): |
thanks |
If this will not be merged then all bounce integrals should be removed per #2225
bounce1doptimization.is_reshaped,is_fourier) that users said were confusing (backwards compatible) as well as the developer flagsBref,Lrefthat should not be there.pitch_batch_sizewas getting ignored. This fixes that by addingstrip_dim0flag tobatch_map.notes