Skip to content

docs(agents): pytree registration classification heuristic - #491

Merged
Jammy2211 merged 1 commit into
mainfrom
docs/jax-pytree-registration-heuristic
Aug 26, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
docs/jax-pytree-registration-heuristic

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Adds §6 — Registering a new type: classification heuristic to docs/agents/jax_and_decorators.md, directly after the §4 Pattern-2 material that already documents register_instance_pytree.

Why

The heuristic lived in PyAutoBrain/skills/register_and_iterate/reference.md. That skill was superseded by /run_queue on 2026-07-08 and has not been invoked since, so it was retired this session — but the pytree knowledge inside it is still current and belongs in the repo that owns the API.

What it documents

  • Classifying an offending class's attributes when a jit trace fails on an unregistered type: all-array/primitive → 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 (the type's wiring site, e.g. _register_fit_imaging_pytrees) and the iterative re-trace habit, with the stop condition — no progress across a few passes means the type holds state that cannot be flattened, which is a design question, not more registrations.
  • The <variant>_pytree.py round-trip assertion used by the parity scripts under autolens_workspace_test/scripts/jax_likelihood_functions/.

Cross-referenced to §5 so it does not contradict the doc it joins: the round trip proves the types flatten, it does not prove xp is threaded — that still needs fitness._vmap(parameters).

Scope

Documentation only — no source, no tests, no API change.

🤖 Generated with Claude Code

https://claude.ai/code/session_01G1SYyCvBWevmYu3x74N8N2

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 <variant>_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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01G1SYyCvBWevmYu3x74N8N2
@Jammy2211
Jammy2211 merged commit da5ec9a into main Aug 26, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the docs/jax-pytree-registration-heuristic branch August 26, 2026 21:03
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