name: jax-python
description: >-
Use for writing, reviewing, debugging, or optimizing Python JAX transformations including jit, grad, value_and_grad, vmap, random keys, pytrees, devices, and array control flow. Do not use for ordinary NumPy code or framework-specific Flax/Equinox/Optax design unless the JAX transformation contract is central.
argument-hint: "[JAX Python task, code, contract, or failure]"
---
# JAX Python
Use the smallest explicit execution contract that preserves semantics. Inspect
the installed version before relying on a drifting signature. State ownership,
input and output shapes, ordering, failure behavior, and the verification command
before writing substantial code.
## Object and execution model
| Object | Meaning | Boundary |
|---|---|---|
| `jax.Array` | A device-backed immutable array value. | Host conversion and device transfer are explicit boundaries. |
| tracer | A symbolic value observed while JAX traces a function. | Python conversion or data-dependent Python control flow is invalid. |
| transformation | jit, grad, vmap, pmap, or another program transform. | The function must be pure over compatible pytrees and shapes. |
| PRNG key | An explicit immutable random-state token. | Split or fold in unique identities; reusing a key repeats randomness. |
| pytree | A nested structure of leaves with registered container shape. | Treedef and leaf shapes/dtypes form part of compiled contracts. |
Read [the object model](references/object-model.md) when the task mixes two
objects or crosses an execution boundary.
## Workflow
1. Write the pure untransformed function and state its pytree, shape, dtype, and randomness contract.
Apply grad/vmap/jit in the minimal required order and mark static choices deliberately.
Keep random keys and updated parameters in explicit return values.
Inspect compilation/retracing and device-transfer boundaries.
Verify eager and transformed outputs plus gradient and reproducibility invariants.
## Decision rules
- Keep transformed functions pure: return updated values and effects as data rather than mutating Python or global state.
Use JAX control-flow primitives for traced data-dependent loops and branches; Python control flow is only for static decisions.
Split keys before independent random uses and fold in stable step/device identities for reproducible parallel work.
Separate static configuration from dynamic arrays and avoid recompilation caused by changing shapes or Python objects.
Keep arrays on device through the compute region; block_until_ready when measuring execution rather than dispatch.
Check gradients against finite differences or a known derivative on small, smooth, nondegenerate inputs.
If a required fact is unknown, inspect the target code, installed signature,
schema, shape, or lifecycle owner. Do not replace an unknown with a permissive
fallback. See [the decision guide](references/decision-guide.md).
## Complex solution routes
Load only the matching section of
[the evaluated recipes](references/recipes-solutions.md):
- `jax.pure-jitted-training-step` and `jax.verify-immutability-and-loss`: compile a pure parameter update without mutating input.
jax.split-fold-vmap-randomness and jax.verify-key-uniqueness-reproducibility: derive reproducible nonreused random streams.
jax.grad-vmap-composition and jax.verify-batch-gradient-shape: vectorize a scalar gradient over a batch.
Recipes are anchors, not blind templates. Preserve their named invariants and
adapt types and names only after inspecting the actual boundary.
## Verification contract
- Test observable behavior, not the presence of API tokens.
- Exercise empty, singleton, malformed, and failure inputs when the operation
accepts them.
- Assert shape, dtype or type, ordering, ownership, and error semantics where
they are part of the contract.
- Keep external I/O deterministic with injected clocks, transports, processes,
files, random state, or test doubles.
- Run the narrow test first, then the relevant project suite. Do not declare
completion when warnings, background failures, convergence flags, or cleanup
errors remain unexplained.
Use [the verification matrix](references/verification.md) for completion checks.
## Failure routing and adaptation
Classify a failure before changing code: input-contract failures require a
precise rejection; environment or version failures require inspection; execution
failures require lifecycle, convergence, or cleanup evidence; invariant failures
require a semantic correction. Do not relax a check, coerce a value, broaden a
failure handler, or materialize data merely to make the symptom disappear.
When adapting a recipe:
1. Match its objects, ownership, execution timing, and output contract to the task.
2. Preserve every branch condition and completion check while changing domain names.
3. Add the project's real empty, malformed, duplicate, cancellation, precision, or
boundary case before removing any guard.
4. If the installed API differs, inspect the signature and primary documentation,
then update implementation, test, and authoring evidence together.
## Version grounding
Inspect the installed package and signature when editing an existing project.
Treat examples here as verified anchors for the version recorded by the Foundry,
not as permission to overwrite a repository's compatibility policy. When current
behavior differs, preserve the project target and update tests and authoring
evidence together.
## Completion
Complete the task only when the implementation preserves the declared object
model, no accidental materialization or lifetime extension was introduced, all
failures are surfaced at the correct boundary, and deterministic tests prove the
critical behavior. Report any environment or version fact that could not be
verified.
1---2name: jax-python3description: ---4---5---6 name: jax-python7 description: >-8 Use for writing, reviewing, debugging, or optimizing Python JAX transformations including jit, grad, value_and_grad, vmap, random keys, pytrees, devices, and array control flow. Do not use for ordinary NumPy code or framework-specific Flax/Equinox/Optax design unless the JAX transformation contract is central.9 argument-hint: "[JAX Python task, code, contract, or failure]"10 ---1112 # JAX Python1314 Use the smallest explicit execution contract that preserves semantics. Inspect15 the installed version before relying on a drifting signature. State ownership,16 input and output shapes, ordering, failure behavior, and the verification command17 before writing substantial code.1819 ## Object and execution model2021 | Object | Meaning | Boundary |22 |---|---|---|23 | `jax.Array` | A device-backed immutable array value. | Host conversion and device transfer are explicit boundaries. |24| `tracer` | A symbolic value observed while JAX traces a function. | Python conversion or data-dependent Python control flow is invalid. |25| `transformation` | jit, grad, vmap, pmap, or another program transform. | The function must be pure over compatible pytrees and shapes. |26| `PRNG key` | An explicit immutable random-state token. | Split or fold in unique identities; reusing a key repeats randomness. |27| `pytree` | A nested structure of leaves with registered container shape. | Treedef and leaf shapes/dtypes form part of compiled contracts. |2829 Read [the object model](references/object-model.md) when the task mixes two30 objects or crosses an execution boundary.3132 ## Workflow3334 1. Write the pure untransformed function and state its pytree, shape, dtype, and randomness contract.352. Apply grad/vmap/jit in the minimal required order and mark static choices deliberately.363. Keep random keys and updated parameters in explicit return values.374. Inspect compilation/retracing and device-transfer boundaries.385. Verify eager and transformed outputs plus gradient and reproducibility invariants.3940 ## Decision rules4142 - Keep transformed functions pure: return updated values and effects as data rather than mutating Python or global state.43- Use JAX control-flow primitives for traced data-dependent loops and branches; Python control flow is only for static decisions.44- Split keys before independent random uses and fold in stable step/device identities for reproducible parallel work.45- Separate static configuration from dynamic arrays and avoid recompilation caused by changing shapes or Python objects.46- Keep arrays on device through the compute region; block_until_ready when measuring execution rather than dispatch.47- Check gradients against finite differences or a known derivative on small, smooth, nondegenerate inputs.4849 If a required fact is unknown, inspect the target code, installed signature,50 schema, shape, or lifecycle owner. Do not replace an unknown with a permissive51 fallback. See [the decision guide](references/decision-guide.md).5253 ## Complex solution routes5455 Load only the matching section of56 [the evaluated recipes](references/recipes-solutions.md):5758 - `jax.pure-jitted-training-step` and `jax.verify-immutability-and-loss`: compile a pure parameter update without mutating input.59- `jax.split-fold-vmap-randomness` and `jax.verify-key-uniqueness-reproducibility`: derive reproducible nonreused random streams.60- `jax.grad-vmap-composition` and `jax.verify-batch-gradient-shape`: vectorize a scalar gradient over a batch.6162 Recipes are anchors, not blind templates. Preserve their named invariants and63 adapt types and names only after inspecting the actual boundary.6465 ## Verification contract6667 - Test observable behavior, not the presence of API tokens.68 - Exercise empty, singleton, malformed, and failure inputs when the operation69 accepts them.70 - Assert shape, dtype or type, ordering, ownership, and error semantics where71 they are part of the contract.72 - Keep external I/O deterministic with injected clocks, transports, processes,73 files, random state, or test doubles.74 - Run the narrow test first, then the relevant project suite. Do not declare75 completion when warnings, background failures, convergence flags, or cleanup76 errors remain unexplained.7778 Use [the verification matrix](references/verification.md) for completion checks.7980 ## Failure routing and adaptation8182 Classify a failure before changing code: input-contract failures require a83 precise rejection; environment or version failures require inspection; execution84 failures require lifecycle, convergence, or cleanup evidence; invariant failures85 require a semantic correction. Do not relax a check, coerce a value, broaden a86 failure handler, or materialize data merely to make the symptom disappear.8788 When adapting a recipe:8990 1. Match its objects, ownership, execution timing, and output contract to the task.91 2. Preserve every branch condition and completion check while changing domain names.92 3. Add the project's real empty, malformed, duplicate, cancellation, precision, or93 boundary case before removing any guard.94 4. If the installed API differs, inspect the signature and primary documentation,95 then update implementation, test, and authoring evidence together.9697 ## Version grounding9899 Inspect the installed package and signature when editing an existing project.100 Treat examples here as verified anchors for the version recorded by the Foundry,101 not as permission to overwrite a repository's compatibility policy. When current102 behavior differs, preserve the project target and update tests and authoring103 evidence together.104105 ## Completion106107 Complete the task only when the implementation preserves the declared object108 model, no accidental materialization or lifetime extension was introduced, all109 failures are surfaced at the correct boundary, and deterministic tests prove the110 critical behavior. Report any environment or version fact that could not be111 verified.