Skip to content

Keep traced initial conditions on device in Tracing.trace - #87

Open
rogeriojorge wants to merge 1 commit into
mainfrom
fix/trace-tracer-initial-conditions
Open

rogeriojorge wants to merge 1 commit into
mainfrom
fix/trace-tracer-initial-conditions

Conversation

@rogeriojorge

Copy link
Copy Markdown
Member

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). Inside jit these arrays are tracers, so the round trip raises jax.errors.TracerArrayConversionError. custom_loss.__call__, grad and value_and_grad are all jitted, so every orbit-based coil objective used through custom_loss is broken, e.g. examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter.py and _fullorbit.py. Plain, non-jitted jax.grad still works.

Fix

A new _place_on_devices(x, target) helper in essos/dynamics.py:

  • Concrete inputs: host placement as before. Typed PRNG keys are still not passed through np.asarray.
  • Tracers, sharded path: lax.with_sharding_constraint(x, NamedSharding).
  • Tracers, single device: passed through unchanged.

Test

tests/test_dynamics.py::test_custom_loss_grad_through_adaptive_guiding_center_matches_finite_difference:

  • Setup: 2 particles traced with GuidingCenterAdaptative through a small BiotSavart coil set (2 base curves, order 1, nfp 2, stellarator symmetric). The loss is taken over the final positions.
  • Check: custom_loss.grad projected on a random direction matches a central finite difference to rtol=1e-6. The observed agreement is about 2e-8 relative.
  • Cost: about 6 s.
  • On main: fails with TracerArrayConversionError.

Checks

  • flake8 error gate (--select=E9,F63,F7,F82): 0.
  • pytest tests: 171 passed, 1 skipped, 1 xfailed.
  • tests/test_dynamics.py with XLA_FLAGS=--xla_force_host_platform_device_count=2, which exercises the sharded branch: 70 passed.

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

codecov Bot commented Sep 27, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 71.42857% with 4 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
essos/dynamics.py 71.42% 3 Missing and 1 partial ⚠️
Files with missing lines Coverage Δ
essos/dynamics.py 72.27% <71.42%> (+0.31%) ⬆️

... and 2 files with indirect coverage changes

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.

1 participant