From 629a0aa299a0ee74201048a0e6a28885890fd2e1 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 26 Aug 2026 16:59:14 -0400 Subject: [PATCH] docs(agents): pytree registration classification heuristic MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds §6 to docs/agents/jax_and_decorators.md, beside the §4 Pattern-2 material that already documents register_instance_pytree: how to classify an offending type's attributes when a jit trace fails on an unregistered class (all-array -> no_flatten=(); the known-aux name set — cosmology, settings, dataset, psf, mask, caches, scipy.spatial.*, Transformer*, PointSolver* —> aux; a callable attribute decided deliberately rather than guessed), where to register it, and the iterative re-trace habit with its stop condition. Carries the _pytree.py round-trip assertion used by the parity scripts in autolens_workspace_test/scripts/jax_likelihood_functions/, and cross-references §5: the round trip proves the types flatten, it does not prove xp is threaded — that still needs fitness._vmap(parameters). Rescued from PyAutoBrain's register_and_iterate skill, retired in that repo this session; the loop mechanics around it were already /run_queue's. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01G1SYyCvBWevmYu3x74N8N2 --- docs/agents/jax_and_decorators.md | 39 +++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/docs/agents/jax_and_decorators.md b/docs/agents/jax_and_decorators.md index 82460b242..4c4cd5fdf 100644 --- a/docs/agents/jax_and_decorators.md +++ b/docs/agents/jax_and_decorators.md @@ -228,3 +228,42 @@ must include **both** a `jax.jit(analysis.fit_from)(instance)` round-trip `autoarray_workspace_test`; array-level JAX changes are exercised downstream in `autogalaxy_workspace_test/scripts/jax_likelihood_functions/` and the `autolens_workspace_test` equivalents. + +--- + +## 6. Registering a new type — classification heuristic + +When a jit trace fails on an unregistered type (a `jax.tree_util` +`TypeError: unhashable type`, a `NotImplementedError`, or a tracer error whose +frame above names the class), read the offending class's `__init__` / +`__dict__` and classify its attributes before registering: + +- **All attrs are `jax.Array` / `np.ndarray` / `AbstractNDArray` / primitives** + → register with `no_flatten=()` (everything dynamic). +- **Known-aux patterns** — `cosmology`, `settings`, `config`, `dataset`, `psf`, + `mask`, `_cache`, `_compiled`, any `scipy.spatial.*`, `Transformer*`, + `PointSolver*` → register with those names in `no_flatten`; they are + per-analysis constants, not traced leaves. +- **A callable attribute** → do not guess. Callable state is the case that + silently breaks gradients; decide aux / dynamic / split deliberately. + +Register at the type's wiring site (e.g. `_register_fit_imaging_pytrees` in +PyAutoLens) with a one-line comment naming the variant that needed it, then +re-trace. Registration is iterative — each pass usually surfaces the next +unregistered type one frame further in — but if a handful of passes produce no +progress, the type is telling you it holds state that cannot be flattened; +stop and treat it as a design question rather than adding more registrations. + +The parity scripts under +`autolens_workspace_test/scripts/jax_likelihood_functions/` are where this is +exercised; `_pytree.py` asserts the round trip against the NumPy +reference: + +```python +ref = analysis.fit_from(instance).log_likelihood # NumPy reference +jit_ll = jax.jit(lambda i: analysis.fit_from(i).log_likelihood)(instance) # jit round-trip +assert jnp.allclose(ref, jit_ll, rtol=1e-4), f"divergence: {ref} vs {jit_ll}" +``` + +That round trip proves the types flatten; it does **not** prove `xp` is threaded +— pair it with the `fitness._vmap(parameters)` check in §5.