Keep traced initial conditions on device in Tracing.trace - #87
Open
rogeriojorge wants to merge 1 commit into
Open
rogeriojorge wants to merge 1 commit into
rogeriojorge wants to merge 1 commit into
Conversation
Tracing.trace placed initial conditions and random keys with a host round trip (device_get + device_put). Inside jit, e.g. custom_loss __call__/grad/value_and_grad, these are tracers and the round trip raised TracerArrayConversionError, breaking every orbit-based coil objective used through custom_loss. Concrete inputs keep the host placement; tracers are constrained with with_sharding_constraint on the sharded path and passed through on a single device. Add a test that custom_loss.grad through a short GuidingCenterAdaptative trace matches a central finite difference.
Codecov Report❌ Patch coverage is
|
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.
Problem
Tracing.trace()placed initial conditions and random keys on devices with a host round trip,device_put(np.asarray(jax.device_get(...)), ...), in both the sharded and single-device branches (introduced in 87bc099 and 3bdf0c6). Insidejitthese arrays are tracers, so the round trip raisesjax.errors.TracerArrayConversionError.custom_loss.__call__,gradandvalue_and_gradare all jitted, so every orbit-based coil objective used throughcustom_lossis broken, e.g.examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter.pyand_fullorbit.py. Plain, non-jittedjax.gradstill works.Fix
A new
_place_on_devices(x, target)helper inessos/dynamics.py:np.asarray.lax.with_sharding_constraint(x, NamedSharding).Test
tests/test_dynamics.py::test_custom_loss_grad_through_adaptive_guiding_center_matches_finite_difference:GuidingCenterAdaptativethrough a smallBiotSavartcoil set (2 base curves, order 1, nfp 2, stellarator symmetric). The loss is taken over the final positions.custom_loss.gradprojected on a random direction matches a central finite difference tortol=1e-6. The observed agreement is about 2e-8 relative.TracerArrayConversionError.Checks
--select=E9,F63,F7,F82): 0.pytest tests: 171 passed, 1 skipped, 1 xfailed.tests/test_dynamics.pywithXLA_FLAGS=--xla_force_host_platform_device_count=2, which exercises the sharded branch: 70 passed.