brainstate.transform.eqns_to_jaxpr#
- brainstate.transform.eqns_to_jaxpr(eqns, invars=None, outvars=None, constvars=None)[source]#
Convert a sequence of JaxprEqn into a Jaxpr.
- Parameters:
eqns (
Sequence[JaxprEqn]) – Sequence of Jaxpr equations to convertinvars (
Sequence[Var]) – Input variables. If None, will be inferred from equationsoutvars (
Sequence[Var]) – Output variables. If None, will be inferred from equationsconstvars (
Sequence[Var]) – Constant variables. If None, will be automatically extracted from equations
- Returns:
A Jaxpr object constructed from the equations
- Return type:
Jaxpr
Notes
constvarsare always placed beforeinvarsin the resulting jaxpr’s binder list, but how they read back differs by jax version. jax 0.11.1 droppedJaxpr’s separate constvar count, so an input now counts as a constvar only when a constant value is attached to it: on jax >= 0.11.1 theconstvarspassed here surface as leadingjaxpr.invarsandjaxpr.constvarsis empty. On jax < 0.11.1 they surface asjaxpr.constvars.Either way the binder order is preserved, so pairing the result with its constant values – which is what
eqns_to_closed_jaxpr()does – yields the sameconstvars/invarssplit on every supported jax version.