brainstate.transform.eqns_to_jaxpr

Contents

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 convert

  • invars (Sequence[Var]) – Input variables. If None, will be inferred from equations

  • outvars (Sequence[Var]) – Output variables. If None, will be inferred from equations

  • constvars (Sequence[Var]) – Constant variables. If None, will be automatically extracted from equations

Returns:

A Jaxpr object constructed from the equations

Return type:

Jaxpr

Notes

constvars are always placed before invars in the resulting jaxpr’s binder list, but how they read back differs by jax version. jax 0.11.1 dropped Jaxpr’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 the constvars passed here surface as leading jaxpr.invars and jaxpr.constvars is empty. On jax < 0.11.1 they surface as jaxpr.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 same constvars/invars split on every supported jax version.