NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #3887 most downloaded on PyPI
Elegant easy-to-use neural networks in JAX.
Last release 5 months ago
05 May 2026
Ships fairly regularly
a new release about every 2 months
Nearly every release is documented
notes for 55 of the last 60 stable releases
3 versions withdrawn
withdrawn after publishing
5 years old
68 releases · first in 2021
One column per quarter.
This is a minor release, largely for internal maintenance purposes.
This is a minor release, largely for internal maintenance purposes.
equinox.internal (=not a public API) (Thanks @ASKabalan! https://github.com/patrick-kidger/equinox/pull/1226)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.7...v0.13.8
JAX v0.10.0 deprecated jax.core.find_top_trace; now using jax.extend.core.find_top_trace when available.
Compatibility:
jax.core.find_top_trace; now using jax.extend.core.find_top_trace when available. (#1218)jax.core.get_aval, jax.core.mapped_aval, jax.core.unmapped_aval; suppressed the resulting type errors. (#1218)equinox.internal.finalise_jaxpr. (Thanks @kasper0406! #1200)Other:
Module.__new__ now returns Self rather than Module. (Thanks @RunnersNum40! #1207)uv and dependency groups for development. (#1218)New awesome-list entries:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.6...v0.13.7
Bug fix: overriding an AbstractClassVar with a concrete instance field will no longer cause a crash when attempting to instantiate the module.
Bug fix: overriding an AbstractClassVar with a concrete instance field will no longer cause a crash when attempting to instantiate the module. (#1198)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.5...v0.13.6
Overriding eqx.Module fields with class attributes should no longer crash when flattening/jitting. (#1191, #1192)
Bugfixes:
eqx.Module fields with class attributes should no longer crash when flattening/jitting. (#1191, #1192)Other
eqx.nn.MLP(..., scan=True), which will use scan-over-layers to improve compilation speed. (Thanks @nstarman! #1179, #1182)eqx.filter_jit(...).lower().compile(...) (Thanks @knyazer! #1178)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.4...v0.13.5
Hotfix release for crash with scan-of-vmap-of-module-with-default-field. (#1172, #1174)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.3...v0.13.4
JAX >=0.8.2 deprecated jax.core.{subjaxprs, mapped_aval, unmapped_aval} (#1157).
EQX_ON_ERROR=off as a valid option. (Thanks @michael-0brien! #1134)jax.core.get_aval -> jax.typeof (#1157).jax.core.{subjaxprs, mapped_aval, unmapped_aval} (#1157).jax.pure_callback in to ValueErrors; we now catch these appropriately (#1156, #1157).eqx.Enumeration.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.2...v0.13.3
More fixes for JAX 0.7.2 compatibility (Thanks @nstarman! #1112)
Compatibility release:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.1...v0.13.2
JAX has deprecated jax.interpreters.batching.{not_mapped, NotMapped}. (Thanks @RunnersNum40! #1086)
The highlights of this release are compatibility with JAX 0.7.2 and with Python 3.14.
Compatibility with JAX 0.7.2.
jax.interpreters.batching.{not_mapped, NotMapped}. (Thanks @RunnersNum40! #1086)jax._src.literals.{LiteralArray, TypedNdArray} (the former is the name in JAX 0.7.2 and the latter is its name in JAX nightly). This now requires handling in all places that a np.ndarray or jax.Array was previously acceptable. (Thanks @johannahaffner! #1099)Compatibility with Python 3.14. (#1100)
cls.__annotations__['__dict__'] is now deprecated in favour of inspect.get_annotations(cls) (https://github.com/python/cpython/issues/139140).eqx.nn.State from being short random strings to instead being readable paths for the location of the state within the pytree. (Thanks @fhchl! #1050, #1052)eqx.field is now compatible with pyright in strict mode. (Thanks @HGangloff! #1053)some_module.foo was being treated as returning Any by static type checkers. Now fixed. (Thanks @nisheethlahoti! #1072)eqx.nn.ConvTranspose not correctly verifying that its stride and output_padding were consistent. (#1045)eqx.nn.ConvTranspose3d ignoring its dtype argument. (Thanks @ZagButNoZig! #1044)eqx.Enumeration.__repr__ crashing when it was the output of a vmap'd function. (Thanks @LennartGevers! #1102)eqx.tree_at documentation improved (Thanks @jeertmans! #1065)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.13.0...v0.13.1
Simplified and tided up the definition of equinox.Module. In principle this should be backward-compatible, but we're doing a minor version bump for th
equinox.Module. In principle this should be backward-compatible, but we're doing a minor version bump for this just in case of edge cases. (#1028)equinox.nn.BatchNorm now has both an ema and a batch mode. For backward compatibility we default to the former, with a warning about choosing explicitly. The latter seems to obtain better performance. (Thanks @lockwo! #659, #948)equinox.{filter,partition} (Thanks @HGangloff! #1038)field(init=False), which can lead to unexpected behaviour. (#1038)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.12.2...v0.13.0
Forward compatibility with the upcoming JAX release, as that removes jaxlib.xla_extension.Device.
jaxlib.xla_extension.Device. (https://github.com/patrick-kidger/equinox/pull/1023)filter_jit boundary. (Thanks @ZagButNoZig! https://github.com/patrick-kidger/equinox/pull/989)doctest. (Thanks @jeertmans! https://github.com/patrick-kidger/equinox/pull/1018)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.12.1...v0.12.2
Hotfix to work around JAX bug https://github.com/jax-ml/jax/issues/27545
Hotfix to work around JAX bug https://github.com/jax-ml/jax/issues/27545 (#988)
Equinox v0.12.0 release notes here: https://github.com/patrick-kidger/equinox/releases/tag/v0.12.0
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.12.0...v0.12.1
BREAKING CHANGE: eqx.field(converter=...) now runs *after* __post_init__ instead of before. This dramatically simplifies some internals and improves c…
eqx.field(converter=...) now runs after __post_init__ instead of before. This dramatically simplifies some internals and improves compatibility with other libraries. (#969, #975)eqx.filter_closure_convert with JAX 0.5.3. (#979, #981)eqx.filter_jit. (Thanks @ZagButNoZig! #973, #980, #983)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.12...v0.12.0
This is primarily a compatibility release.
This is primarily a compatibility release.
eqx.nn.Linear(0, ...) no longer crashes runs (propagating zero-size tensors). (Thanks @aseyboldt! #950)eqx.tree_pformat and eqx.tree_pprint) now uses the new Wadler-Lindig pretty-printing library. (#924)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.11...v0.11.12
JAX 0.4.38 moved a number of APIs with a deprecation warning, e.g. jax.core.Jaxpr -> jax.extend.core.Jaxpr. With this release we've updated and are ba…
JAX 0.4.38 moved a number of APIs with a deprecation warning, e.g. jax.core.Jaxpr -> jax.extend.core.Jaxpr. With this release we've updated and are back to being warning-free under this JAX release! (Thanks @FFroehlich @DrJessop, #913, #915, #917)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.10...v0.11.11
This is a JAX 0.4.36 compatibility release.
This is a JAX 0.4.36 compatibility release.
With this release, JAX changed how custom primitive rules are called (they are always called, instead of only when the data requires them to be). That requires some updates in Equinox to avoid crashes in the downstream ecosystem. (https://github.com/patrick-kidger/diffrax/issues/532, https://github.com/jax-ml/jax/issues/25289 + links therein.)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.9...v0.11.10
This is a (important) bugfix release.
This is a (important) bugfix release.
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.8...v0.11.9
The main thing for this release is JAX 0.4.34 compatibility -- JAX introduced breaking changes in this release that we are now compatible with.
The main thing for this release is JAX 0.4.34 compatibility -- JAX introduced breaking changes in this release that we are now compatible with. (#871)
__init_subclass__ should no longer crash. (Plus probably better-behaved __init_subclass__ overall.)eqx.error_if's nice displaying of error message. With this release then we are back to having nice error messages again!eqx.nn.StateIndex can now be passed through jax.jit (and not just eqx.filter_jit). (Thanks @NeilGirdhar! #843)~= version constraints. We now work around that for better compatibility with certain kinds of Poetry installations. (Thanks @norpadon! #878)eqx.tree_at documentation for clarity. (Thanks @jeertmans! #872, #874, #877)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.7...v0.11.8
Quick release. JAX 0.4.32 / 0.4.33 just introduced a breaking change; this release ensures Equinox is compatible with this.
Quick release. JAX 0.4.32 / 0.4.33 just introduced a breaking change; this release ensures Equinox is compatible with this. (#856)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.6...v0.11.7
This is primarily a bug fix release.
This is primarily a bug fix release.
Runtime error messages (those from eqx.error_if, in particular when wrapped with eqx.filter_jit) should now be compatible with PyCharm's debugger, and with certain multithreaded contexts. (Thanks @adam-hartshorne, @dlwh! #828, #844, #849)
Marking a jax.Array or np.ndarray as an eqx.field(static=True) will now raise a warning. This was technically okay as long as you use it in certain very narrow contexts (e.g. to smuggle it into a JIT'd region without being traced), but in practice it was nearly always just a common new-user footgun. (Thanks @lockwo! #800)
Using eqx.tree_at for replacing empty tuples is improved. (Thanks @danielward27! #818, #819)
eqx.nn.RotaryEmbedding no longer promote input dtypes to at least float32. (Thanks @knyazer! #836)
Mypy now understands that eqx.Modules are dataclasses. (Pyright always did, but mypy needed a slightly different approach to appreciate this fact.) (Thanks @NeilGirdhar! #822)
Multiple eqx.Modules participating in co-operative multiple inheritance (at least 5 inheriting from each other seem to be necessary?), with some of them overriding the __post_init__s of others, should now follow their expected resolution order. (Thanks @NeilGirdhar! #832, #834)
We now have a .editorconfig file, (thanks @NeilGirdhar! #821)
Doc improvements. (Thanks @garymm, @ColCarroll! #804, #805)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.5...v0.11.6
Recent versions of JAX (0.4.28+) have made some changes to:
Recent versions of JAX (0.4.28+) have made some changes to:
With this update, we should now be compatible with both old and new versions of JAX: this fixes both some new crashes, and some new warnings. (#719, #724, #753, #758, thanks @jakevdp, @hawkinsp!)
The error messages from eqx.error_if are now substantially more informative: they include traceback information including the stack, and mention the availability of the EQX_ON_ERROR variable. We also do a much better job hiding the large unhelpful printouts that XLA gives by default. (#785, #803)
The default value of EQX_ON_ERROR_BREAKPOINT_FRAMES is now 1. (#777) The impact of this is that using eqx.error_if alongside EQX_ON_ERROR=breakpoint will now:
u).eqx.filter_{jacfwd, jacrev} now only apply filtering to their inputs but not their outputs. Previously this was problematic as there was no way to represent static-input-by-static-output in the returned Jacobian, so pieces were silently dropped. (#734, thanks @lockwo!)
eqx.tree_at can now be used to replace empty tuples. (#715, #717, #722, thanks @lockwo!)
eqx.filter_custom_jvp no longer raises a trace-time crash in some scenarios in which its **kwargs were erroneously counted as having tangents. (https://github.com/patrick-kidger/equinox/issues/745#issuecomment-2148560546, #749)
No longer getting a trace-time crash when doing a particular combination of vmap + autodiff + checkpointed while loops. This occurred when using optimistix.BFGS around diffrax.diffeqsolve. (#777)
Fixed a trace-time crash when:
Fixed a trace-time crash when composing the grad of vmap of lineax.linear_solve. (https://github.com/patrick-kidger/lineax/issues/101, #795, thanks @rhacking!)
eqx.nn.RMSNorm now uses at least 32-bit precision for numerical stability (#723, thanks @AakashKumarNain!)
eqx.nn.{Linear,Conv,GRUCell,LSTMCell} now support complex dtypes (#765, thanks @ChenAo-Phys!)
Added eqx.nn.RotaryEmbedding(..., theta=...). (#735, thanks @Artur-Galstyan!)
Several doc fixes. (#708, #731, #733, #747, #750, #757 + several other PRs, thanks @Artur-Galstyan, @matteoguarrera, @lockwo, @nasyxx!)
Several internal test fixes as downstream libraries have changed slightly. (#740, #742 + several other PRs, big thanks to @GaetanLepage for reporting many of these!)
There is now a Mistral 7B implementation using JAX+Equinox available over in AakashKumarNain/mistral_jax!
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.4...v0.11.5
Added eqx.filter_shard. This lowers to jax.lax.with_sharding_constraint as a single way to transfer data, or reshard data, both inside and outside of
Added eqx.filter_shard. This lowers to jax.lax.with_sharding_constraint as a single way to transfer data, or reshard data, both inside and outside of JIT! (No more jax.device_put.) In addition, the parallelism example has been updated to use this simpler new functionality. (Thanks @homerjed and @dlwh! #688, #691)
Added eqx.filter_{jacfwd,jacrev,hessian}. These do what you expect! (Thanks @lockwo! #677)
Added eqx.nn.RotaryPostionalEmbedding. This is designed to be used in conjunction with the existing eqx.nn.MultiheadAttention. (Thanks @Artur-Galstyan! #568)
Added support for padding='VALID', padding='SAME', padding='SAME_LOWER' to the convolutional layers: eqx.nn.{Conv, ...}. (Thanks @ChenAo-Phys! #658)
Added support for padding_mode='ZEROS', padding_mode='REFLECT', padding_mode='REPLICATE', padding_mode='CIRCULAR' to the convolutional layers: eqx.nn.{Conv, ...}. (Thanks @ChenAo-Phys! #658)
Added a dtype argument to eqx.nn.{MultiheadAttention, Linear, Conv, ...} for specifying the dtype of their parameters. In addition eqx.nn.BatchNorm will now also uses its dtype argument to determine the dtype of its weights and bias, not just the dtype of its moving statistics. (Thanks @Artur-Galstyan and @AakashKumarNain! #680, #689)
eqx.error_if is now compatible with JAX 0.4.26, which changed JAX's own reporting of error messages slightly. (Thanks @hawkinsp! #670)
Added a warning that checks for doing something like:
class MyModule(eqx.Module):
fn: Callable
def __init__(self, ...):
self.fn = jax.vmap(some_fn)
As this is an easy source of bugs. (The vmap'd function is not a PyTree so will not propagate anything in the PyTree stucture of some_fn.)
eqx.internal.while_loop(..., kind="checkpointed") will now only propagate forward JVP tracers for those outputs which are perturbed due to the input to the loop being perturbed. (Rather than all of them.) This change just means that later calls to a nondifferentiable operation, like jax.pure_callback or eqx.internal.nondifferentiable, will no longer crash at trace time. (See https://github.com/patrick-kidger/diffrax/issues/396.)
eqx.internal.while_loop(..., kind="bounded") will now handle certain vmap+grad combinations without crashing. (It seems like JAX is adding some spurious batch tracers.) (See https://github.com/patrick-kidger/optimistix/issues/48#issuecomment-2009221739)
the transpose rule for eqx.internal.create_vprim now understands symbolic zeros, fixing a crash for grad-of-vmap-of-<lineax.linear_solve that we only use some outputs from>. (See https://github.com/patrick-kidger/optimistix/issues/48.)
The type annotation for the input of any converter function used in eqx.field(converter=...) will now be used as the type annotation in any dataclass-autogenerated __init__ functions. In particular this should mean such functions are now compatible with runtime type checkers like beartype. (jaxtyping users, you were already covered: this checks the assigned annotations instead.)
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.3...v0.11.4
equinox.tree_deserialise_leaves now treats jax.ShapeDtypeStructs in the same way as arrays. This makes it possible to avoid instantiating the initial
equinox.nn.RMSNorm.equinox.nn.WeightNorm.equinox.tree_deserialise_leaves now treats jax.ShapeDtypeStructs in the same way as arrays. This makes it possible to avoid instantiating the initial model parameters only to throw them away again, by using equinox.filter_eval_shape:model = eqx.filter_eval_shape(Model, ...hyperparameters...)
model = eqx.tree_deserialise_leaves(load_path, model)
(#259)equinox.internal.noinline no longer initialises the JAX backend on use.equinox.filter_jit(...).lower(..., some_kwarg=...) no longer crashes (#625, #627)equionx.nn.BatchNorm now uses the default floating point dtype, rather than always using float32.equinox.nn.MultiheadAttention should now perform the softmax in float32 even when the input is of lower dtype. (This is important for numerical stability.)equinox.nn.{Linear, MLP, ...} now standardise on accepting extra **kwargs and not calling super().__init__. The intention is that these layers be treated as final, i.e. not subclassable. (Previously things were inconsistent: some did this and some did not.)JAX_NUMPY_DTYPE_PROMOTION=strict and JAX_NUMPY_RANK_PROMOTION=raise, and this is checked in tests.filter_grad (Thanks @knyazer! #589)These are undocumented internal features, that may be changed at any time.
EQX_GETKEY_SEED for use with equinox.internal.GetKey.equinox.internal.while_loop now has its runtime errors removed. This should help with compatibility with TPUs. (#628)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.2...v0.11.3
Added eqx.filter_jit(..., donate="all-except-first") and eqx.filter_jit(..., donate="warn-except-first"). This offers a way to donate all arguments *e
Features
eqx.filter_jit(..., donate="all-except-first") and eqx.filter_jit(..., donate="warn-except-first"). This offers a way to donate all arguments except the first one. (If you have multiple such arguments then just pack them together into a tuple in the first argument.) This aims to be a low-overhead easy way to handle buffer donation.eqx.debug.{assert_max_traces, get_num_traces}, which aim to provide a friendly way of asserting that a JIT'd function is not recompiled -- and if it is, which argument changed to cause the recompilation.eqx.tree_pprint and eqx.tree_pformat now handle PyTorch tensors and jax.ShapeDtypeStructs.eqx.tree_equal now has new arguments:
typematch=True: this will require that every leaf have precisely the same type as each other, i.e. right now the requirement is essentially leaf == leaf2; with this flag it becomes type(leaf) == type(leaf2) and leaf == leaf2.rtol and atol: setting these to nonzero values allows for checking that inexact (floating or complex) arrays are allclose, rather than exactly equal.assert eqx.tree_equal(output, expected_output, typematch=True, rtol=1e-5, atol=1e-5).Bugfixes
eqx.nn.MLP would use the exact same learnt weights for every neuron in every layer. Now, a separate copy of the activation function is used in each location.eqx.Module should now have their __init__ signatures correctly reported by downstream tooling, e.g. automated doc generators, some IDEs. (Thanks @danielward27! #573)Typing
eqx.filter_value_and_grad now declares that it preserves the return type of its function (Thanks @ConnorBaker! #557)Documentation
StateIndex (Thanks @edwardwli! #556)eqx.Enumueration docstrings (Thanks @LouisDesdoigts! #579)Other
__tracebackhide__ = True assignments.eqx.Enumerations can now override the message associated with their parent Enumeration: this now produces a warning rather than an error.EQX_ON_ERROR_BREAKPOINT_FRAMES config variable, which is used to work around a JAX bug when setting EQX_ON_ERROR=breakpoint.eqx.Module, e.g.class Foo(eqx.Module):
def f(self): ...
Foo.f = some_transform(Foo.f)
the anticipated use-case for this is to make it easier for typecheckers; see #584.eqx.debug.store_dce now supports non-arrays in its argument.eqx.Enumeration.where(traced_pred, x, x) will now statically return x without tracing. This is occasionally useful to better propagate information at compile time.Internal features (not officially supported, advanced use only)
eqx.internal.GetKey. This generates a random JAX PRNG key when called, and crucially has a nice __repr__ reporting what the seed value is. This should not be used in normal JAX code! This is intended as a convenience for tests, so that the random seed appears in the debug printout of a failed test.eqx.internal.MaybeBuffer to indicate that an argument of an eqx.internal.{while_loop,scan} might be wrapped in a buffer.eqx.internal.buffer_at_set to support buffer.at[...].set(..., pred=...) whilst being agnostic to whether buffer is a JAX array or one of our while loop buffers.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.1...v0.11.2
This is a minor bugfix release.
This is a minor bugfix release.
Bugfixes
eqx.internal.while_loop(..., kind="checkpointed")) now perform a more careful analysis of which arguments need to be differentiated. (#548) This fix is the primary reason for this release -- it unlocks some efficiency improvements when solving SDEs in Diffrax: https://github.com/patrick-kidger/diffrax/pull/320Abstract{Class,}Var misbehaving around multiple inheritance. (#544)Documentation
Other
py.typed marker file. Thanks @vidhanio! #547)EQX_ON_ERROR_BREAKPOINT_FRAMES environment variable, to work around JAX bug https://github.com/google/jax/issues/16732 when using EQX_ON_ERROR=breakpoint. This new variable sets the number of stack frames you can access via the u debugger command, when the on-error debugger is triggered. Set this to a small enough number, e.g. EQX_ON_ERROR_BREAKPOINT_FRAMES=1, and it should fix unusual trace-time errors when using EQX_ON_ERROR=breakpoint.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.11.0...v0.11.1
Breaking change: eqx.nn.StateIndex now takes the initial value, rather than a function that returns the initial value.
Equinox now includes several additional checks to guard against various bugs. If you have a new error, then this is probably an indication that your code always had a silent bug, and should be updated.
eqx.nn.LayerNorm now correctly validates that the shape of its input. This was a common cause of silent bugs. (Thanks @dlwh for pointing this one out!)__init__ and __post_init__ -- the former actually overwrites the latter. (This is normal Python dataclass behaviour, but probably unexpected.)class Model(eqx.Module):
foo: Callable
def __init__(self):
self.foo = self.bar
def bar(self):
...
Otherwise, you end up with two different copies of your model! One at self, the other at self.foo.__self__. (The latter being in the bound method.)eqx.tree_at now gives a better error message if you use it try to and update something that isn't a PyTree leaf. (Thanks @LouisDesdoigts!)These should all be very minor.
eqx.nn.StateIndex now takes the initial value, rather than a function that returns the initial value.eqx.field(converter=...), then conversion now happens before __post_init__, rather than after it.eqx.nn.make_with_state over eqx.nn.State. The latter will continue to work, but the former is more memory-efficient. (It deletes the original copy of the initial state.)eqx.nn.inference_mode over eqx.tree_inference. The latter will continue to exist for backward compatibility. These are the same function, this is really just a matter of moving it into the eqx.nn namespace where it always belonged.Equinox now supports sharing a layer between multiple parts of your model! This has probably been our longest-requested feature -- in large part because of how intractable it seemed. Equinox models are PyTrees, not PyDAGs, so how exactly are we supposed to have two different parts of our model point at the same layer?
The answer turned out to be the following -- in this example, we're reusing the embedding weight matrix between the initial embedding layer, and the final readout layer, of a language model.
class LanguageModel(eqx.Module):
shared: eqx.nn.Shared
def __init__(self):
embedding = eqx.nn.Embedding(...)
linear = eqx.nn.Linear(...)
# These two weights will now be tied together.
where = lambda embed_and_lin: embed_and_lin[1].weight
get = lambda embed_and_lin: embed_and_lin[0].weight
self.shared = eqx.nn.Shared((embedding, linear), where, get)
def __call__(self, tokens):
# Expand back out so we can evaluate these layers.
embedding, linear = self.shared()
assert embedding.weight is linear.weight # same parameter!
# Now go ahead and evaluate your language model.
...
here, eqx.nn.Shared(...) simply removes all of the nodes at where, so that we don't have two separate copies. Then when it is called at self.shared(), it puts them back again. Note that this isn't a copy and doesn't incur any additional memory overhead; this all happens at the Python level, not the XLA level.
(The curious may like to take a look at the implementation in equinox/nn/_shared.py, which turned out to be very simple.)
On a meta level, I'd like to comment that I'm quite proud of having gotten this one in! It means that Equinox now supports both stateful layers and shared layers, which have always been the two pieces that seemed out of reach when using something as simple as PyTrees to represent models. But it turns out that PyTrees really are all you need. :D
eqx.filter_checkpoint, which as you might expect is a filtered version of jax.checkpoint. (Thanks @dlwh!)eqx.Module.__check_init__. This is run in a similar fashion to __post_init__; see the documentation. This can be used to check that invariants of your module hold after initialisation.eqx.nn.State.{substate, update}. This offers a way to subset or update a State object, that so only the parts of it that need to be vmap'd are passed in. See the stateful documentation for an example of how to do this.INTERNAL: Generated function failed: CpuCallback error stuff! This clean-up of the runtime error message is done by eqx.filter_jit, so that will need to be your top-level way of JIT'ing your computation.eqx.nn.StatefulLayer -- this is (only!) with eqx.nn.Sequential, to indicate that the layer should be called with x, state, and not just x. If you would like a custom stateful layer to be compatible with Sequential then go ahead and subclass this, and potentially implement the is_stateful method. (Thanks @paganpasta!)eqx.nn.* layer is now wrapped in a jax.named_scope, for better debugging experience. (Thanks @ahmed-alllam!)eqx.module_update_wrapper no longer requires a second argument; it will look at the __wrapped__ attribute of its first argument.eqx.internal.closure_to_pytree, for... you guessed it, turning function closures into PyTrees. The closed-over variables are treated as the subnodes in the PyTree. This will operate recursively so that closed-over closures will themselves become PyTrees, etc. Note that closed-over global variables are not included.eqx.tree_{serialise,deserialise}_leaves now correctly handle unusual NumPy scalars, like bfloat16. (Thanks @colehaus!)eqx.field(metadata=...) arguments no longer results in the static/converter arguments being ignored. (Thanks @mjo22!)eqx.filter_custom_vjp now supports residuals that are not arrays. (The residuals are the pytree that is passed between the forward and backward pass.)eqx.{AbstractVar,AbstractClassVar} should now support overriden generics in subclasses. That is, something like this:class Foo(eqx.Module):
x: eqx.AbstractVar[list[str]]
class Bar(Foo):
x: list[str]
should no longer raise spurious errors under certain conditions.eqx.internal.while_loop now supports using custom (non-Equinox) pytrees in the state.eqx.tree_check no longer raises some false positives.__init_subclass__ with additional class creation kwargs. (Thanks @ASEM000, @Roger-luo!)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.11...v0.11.0
Equinox now offers true runtime errors! This is available as equinox.error_if. This is something new under the JAX sun: these are raised eagerly durin
Equinox now offers true runtime errors! This is available as equinox.error_if. This is something new under the JAX sun: these are raised eagerly during the execution, they work on TPU, and if you set the environment variable EQX_ON_ERROR=breakpoint, then they'll even drop you into a debugger as soon as you hit an error. (These are basically a strict improvement over jax.experimental.checkify, which doesn't offer many of these advantages.)
Added a suite of debugging tools:
equinox.debug.announce_transform: prints to stdout when it is transformed via jvp/vmap etc; very useful for keeping track of how many times a particular operation is getting transformed or compiled, when trying to minimise your compilation times.equinox.debug.backward_nan: for debugging NaNs that only arise on the backward pass.equinox.debug.breakpoint_if: opens a breakpoint if a condition is satisfied.equinox.debug.{store_dce, inspect_dce}: used for checking whether certain variables are removed via the dead-code-elimination pass of the XLA compiler.equinox.filter_jvp now supports keyword arguments (which are treated as not differentiated).
filter_jvps will now no longer materialise symbolic zero tangents. (#422).Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.10...v0.10.11
These are the real highlight of this release.
Performance improvements
These are the real highlight of this release.
equinox.internal.{while_loop, scan} now use new symbolic zero functionality, which may result in runtime speedups (and slight increases in compile times) as they can now skip calculating gradients for some quantities.equinox.internal.{while_loop, scan}(..., buffers=...) now do their best to work around an XLA bug (https://github.com/google/jax/issues/10197). This can reduce computational cost from quadratic scaling to linear scaling.equinox.internal.{while_loop, scan} now includes several optimisations for the common case is which every step is checkpointed. (#415)Features
equinox.filter_custom_{jvp,vjp} now support symbolic zeros.
Previously, None was passed to represent symbolic zero tangent/cotangents for anything that wasn't a floating-point array -- but all floating-point-arrays always had materialised tangent/cotangents.
With this release, None may also sometimes be passed as the tangent of floating-point arrays. In this case it represents a zero tangent/cotangent, and moreover this zero is "symbolic" -- that is to say it is known to be zero at compile time, which may allow you to write more-efficient custom JVP/VJP rules. (The canonical example is the inverse function theorem -- this involves a linear solve, parts of which you can skip if you know parts of it are zero.)
In addition, filter_custom_vjp now takes another argument, perturbed, indicating whether a value actually needs cotangents calculated for it. You can skip calculating cotangents for anything that is not perturbed.
For more information see jax.custom_jvp.defjvp(..., symbolic_zeros=True) and jax.custom_vjp.defvjp(..., symbolic_zeros=True), which provide the underlying behaviour that is being forwarded.
Note that this is provided through a new API: filter_custom_jvp.def_jvp instead of filter_custom_jvp.defjvp, and filter_custom_vjp.{def_fwd, def_bwd} instead of filter_custom_vjp.defvjp. The old API will continue to exhibit the previous behaviour, for backward compatibility.
Misc
New Contributors
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.6...v0.10.10
(Why no v0.10.{7,8,9}? We had a bit of a rocky release this time around, and these got yanked for having bugs. Thanks to everyone who reported issues so quickly! Things look like they're stable now...)
Nothing published for this version
Nothing published for this version
Nothing published for this version
Added eqx.field: this supports converter=... and static=.... The former is an extension to dataclasses that applies that conversion function when the
eqx.field: this supports converter=... and static=.... The former is an extension to dataclasses that applies that conversion function when the field is assigned. The latter supersedes the old eqx.static_field. (#390)eqx.Enumeration, which are JAX-compatible Enums. (Moved from `eqx.internal.Enumeration.) (#392)eqx.clear_caches to clear internal caches and reduce memory usage. (#380)eqx.nn.BatchNorm(..., dtype=...) (Thanks @Benjamin-Walker! #384)eqx.internal.while_loop: buffers now support buffer.at[index].add(...) etc. (Thanks @packquickly! #395)typing->collections.abc where appropriate; Tuple->tuple etc. (#385)eqx.module_update_wrapper no longer assigns __wrapped__. (#381)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.5...v0.10.6
Fixed modules initialising twice (#369; this bug was introduced in the last couple of Equinox versions.)
Quite a small release.
Bugfixes
Documentation
MLP.__init__. (Thanks @schmrlng! #366)Misc
equinox.internal.eval_full (like equinox.internal.eval_{zeros, empty}) (Thanks @RaderJason! #367)equinox.internal.Enumeration (#375)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.4...v0.10.5
eqx.nn.{LayerNorm, GroupNorm} can now accept a call-time state argument that they thread through unchanged. This means that they have the same API as
eqx.nn.{LayerNorm, GroupNorm} can now accept a call-time state argument that they thread through unchanged. This means that they have the same API as eqx.nn.BatchNorm, so that they may be used interchangeably.eqx.Modules now work with the new jax.tree_util.tree_flatten_with_path API. (#363)eqx.nn.MLP now supports use_bias and use_final_bias arguments. (Thanks @jlperla! #358)eqx.tree_check to assert that a pytree does not contain duplicate elements, and does not contain any reference cycles. This may be useful to call on your models prior to training, to check that they are well-formed. (#355)eqx.tree_flatten_one_level to flatten a pytree by one level only. (#355)eqx.internal.{error_if, branched_error_if, debug_backward_nans} now have TPU support! This means that they now support all backends, and are (to my knowledge) the single best option for adding runtime checks to JAX programs. In addition they now eagerly will raise errors at trace-time if the predicate is a raw Python True. (#351)eqx.internal.scan now supports buffers and checkpoints arguments for finer-grained control over its autodiff. (#349)eqx.internal.scan_trick, which can be used to minimise compilation time by wrapping nearby function invocations into a single scan. See this PR against Diffrax for an example.eqx.nn.ConvTranspose (Thanks @khdlr! #335)eqx.static_field()s were sometimes being put in leaves; this is now fixed. (This issue existed in v0.10.3 only.) (#338)eqx.filter_custom_jvp will no longer raise the occasional spurious leaked tracer error. (When using traced non-floating arrays.) (#349)eqxi.while_loop(... kind='checkpointed') (#331)pyproject.toml to handle everything (no more setup.py, .flake8 etc!)eqx.internal.while_loop should now have a slightly faster compile time. (#353)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.3...v0.10.4
Added equinox.nn.{State, StateIndex}. This has been one of the longest-requested features for Equinox: we now have proper stateful operations! (In a c
Added equinox.nn.{State, StateIndex}. This has been one of the longest-requested features for Equinox: we now have proper stateful operations! (In a carefully-controlled way -- see the new stateful docs and the new stateful example.)
As an application of these new stateful operations: added equinox.nn.{BatchNorm, SpectralNorm}, which have graduated from experimental! Note that these have a slightly different API to their previous experimental versions.
Added equinox.Partial, which is a tidied-up version of jax.tree_util.Partial.
equinox.filter_{jit, pmap} are now compatibile with ahead-of-time compilation. (#325)
equinox.nn.LayerNorm now supports use_weight and use_bias arguments to disable each individually. This is reflecting the fact that many modern transformer architectures now use layer normalisation without bias. (#310; thanks @lockwo!)
Added equinox.internal.{AbstractVar, AbstractClassVar} to denote abstract instance attributes and abstract class attributes respectively. (Analogous to abc.abstractmethod denoting abstract methods.) The downstream scientific ecosystem is making heavy use of abstract base classes (e.g. all the ABCs in Diffrax) and these have turned out to be a really useful feature. See this docstring for more details. Right now these are an undocumented internal-only feature, but we could plausibly spin these out into their own library.
equinox.nn.Conv should now be compatible with disabled rank promotion (#308; thanks @lockwo!)equinox.internal.loop should now be compatible with jax.experimental.xmap (https://github.com/patrick-kidger/diffrax/issues/246)equinox.experimental.* has been removed. See the new stateful functionality described above.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.2...v0.10.3
Fixed (the deprecated, but still) deterministic=True being ignored in MultiheadAttention. (Thanks @mk-0!)
This release has lots of examples and bugfixes from several new contributors!
Features
eqx.nn.{Linear, MLP} now support the string "scalar" for their input and output sizes, to produce an array of shape () rather than an array of shape (1,).equinox.internal.scan for a checkpointed scan implementation. (It'd be interesting to see this used for an optimally-checkpointed scan-over-layers in an LLM?)Documentation
Bugfixes
eqx.filter_closure_convert and eqx.internal.while_loop now work with tree-math.MultiheadAttention, and fixed it producing NaNs in fully-masked case. (Thanks @j5b!)deterministic=True being ignored in MultiheadAttention. (Thanks @mk-0!)__new__ can now be overridden in subclasses of eqx.Module. (Thanks @ASEM000!)Misc
equinox._jit instead of equinox.jit. If you're broken by this change then you should make sure to import from the public interface: e.g. equinox.filter_jit instead of equinox._jit.filter_jit.equinox.internal.while_loop(..., kind="checkpointed") now supports readable buffers.eqx.filter_vmap now supports all-Nones in in_axes. (Thanks @RaderJason!)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.1...v0.10.2
See the v0.10.0 release notes for the interesting recent changes.
The usual post-release hotfix.
See the v0.10.0 release notes for the interesting recent changes.
Changes in this release
typing_extensions.eqx.filter_{vmap, pmap}(in_axes=dict(...), ...) crashing when used alongside default arguments.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.10.0...v0.10.1
A dramatically simplified API for equinox.{filter_jit, filter_grad, filter_value_and_grad, filter_vmap, filter_pmap} . This is a backward-incompatible
A dramatically simplified API for equinox.{filter_jit, filter_grad, filter_value_and_grad, filter_vmap, filter_pmap} . This is a backward-incompatible change.
equinox.internal.while_loop, which is a reverse-mode autodifferentiable while loop, using recursive checkpointing.
Some new relatively minor new features available in this release.
eqx.{filter_jit, filter_pmap}. (Thanks @uuirs in #235!)eqx.nn.PRelu. (Thanks @enver1323 in #249!)eqx.tree_pprint.eqx.module_update_wrapper.eqx.filter_custom_jvp now supports keyword arguments (which are always treated as nondifferentiable).internal featuresIntroducing a slew of new features for the advanced JAX user.
These are all available in the equinox.internal namespace. Note that these comes without stability guarantees, as they often depend on functionality that JAX doesn't make fully public.
eqxi.abstractattribute, for marking abstract instance attributes of abstract Equinox modules.eqxi.tree_pp, for producing a pretty-print doc of an object. (This is what is then formatted to a particular width in e.g. eqx.tree_pformat.) In addition classes can now have custom pretty behaviour when used with eqx.{tree_pp, tree_pformat, tree_pprint}, by setting a __tree_pp__ method.eqxi.if_mapped, as an alternative to the usual eqx.if_array passed to eqx.{filter_vmap, filter_pmap}(out_axes=...).eqxi.{finalise_jaxpr, finalise_fn} for tracing through custom primitives impl rules (so that the custom primitive no longer appears in the jaxpr). This is useful for replacing such custom primitives prior to offloading a jaxpr to some other IR, e.g. via jax2tf.eqxi.{nonbatchable, nondifferentiable, nondifferentiable_backward, nontraceable} for asserting that an operation is never batched, differentiated, or subject to any transform at all.eqxi.to_onnx for exporting to ONNX.eqxi.while_loop for reverse-mode autodifferentiable while loops; in particular making use of recursive checkpointing. (A la treeverse.)equinox.{filter_jit, filter_grad, filter_value_and_grad, filter_vmap, filter_pmap} has been dramatically simplified. If you were using the extra arguments to these functions (i.e. not just calling @eqx.filter_jit etc. directly) then this is a backward-incompatible change; see the discussion below for more details.equinox.nn.{AvgPool1D, AvgPool2D, AvgPool3D, MaxPool1D, MaxPool2D, MaxPool3D}. Use AvgPool1d etc. (lower-case "d") instead. (These were backward-compatiblity stubs that have now been removed.)equinox.Module.{tree_flatten, tree_unflatten}. These were never technically public API; use jax.tree_util.{tree_flatten, tree_unflatten} instead.equinox.filter_closure_convert now asserts that you call it with argments compatible with those it was closure-converted with.filter_jit or filter_pmap boundary should now be much reduced.eqx.tree_inference now runs faster. (Thanks @uuirs in #233!)These APIs have been simplified and made much easier to understand. No functionality has been lost, things might just need tweaking.
filter_jitThis previously took default, args, kwargs, out, fn arguments, for controlling what should be traced and what should be held static.
In practice all JAX arrays and NumPy arrays always had to be traced, and everything that wasn't a JAXable type (JAX array, NumPy array, bool, int, float, complex) had to be held static. So these arguments just weren't that useful: pretty much the only thing you could do with them was to specify that you'd like to trace a bool/int/float/complex.
This minor use-case wasn't worth complicating such an important API for, which is why these arguments have been removed.
If after this change you still want to trace with respect to bool/int/float/complex, then do so simply by wrapping them into JAX arrays or NumPy arrays first: np.asarray(x).
filter_grad and filter_value_and_gradThese previously took an arg argument, for controlling what parts of the first argument should be differentiated.
This was useful occasionally -- e.g. when freezing parts of a layer -- but in practice it still wasn't used that often. As such it this argument has been removed for the sake of simplicity.
If after this change you want to replicate the previous behaviour, then it is simple to do so using partition and combine:
# Before
@eqx.filter_grad(arg=foo)
def loss(first_arg, ...):
...
loss(bar, ...)
# After
@eqx.filter_grad
def loss(diff_first_arg, static_first_arg, ...):
first_arg = eqx.combine(diff_first_arg, static_first_arg)
...
diff_bar, static_bar = eqx.partition(bar, foo)
loss(diff_bar, static_bar, ...)
See also the updated frozen layer example for a demonstration.
filter_vmapThis previously took default, args, kwargs, out, fn arguments, for controlling what axes should be vectorised over.
In practice this API was just a bit more complicated than it really needed to be. The only useful feature relative to jax.vmap was kwargs, for easily specifying just a few named arguments that should behave differently.
The new API instead accepts in_axes and out_axes arguments, just like jax.vmap. To replace kwargs, one extra feature is supported: in_axes may be a dictionary of named argments, e.g.
@eqx.filter_vmap(in_axes=dict(bar=None))
def fn(foo, bar):
...
All arguments not named in kwargs will have the default value of eqx.if_array(0) -> 0 if is_array(x) else None applied to them.
On which note, a new eqx.if_array(i) now exists, to make it easier to specify values for in_axes and out_axes.
If you were using the old fn argument, then this can be replicated by instead decorating a function that accepts the callable:
# Before
@eqx.filter_vmap(foo, fn=bar)(x, y)
# After
@eqx.filter_vmap(in_axes=dict(fn=bar))
def accepts_foo(fn, x, y):
return fn(x, y)
accepts_foo(foo, x, y)
filter_pmap.This previously took default, args, kwargs, out, fn arguments, for controlling what axes should be parallelised over, and which arguments should be traced vs static.
This was a fiendishly complicated API merging together both the filter_jit and filter_vmap APIs.
The JIT part of it is now handled automatically, as with filter_jit: all arrays are traced, everything else is static.
The vmap part of it is now handled in the same way as filter_vmap, using in_axes and out_axes arguments.
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.9.2...v0.10.0
Autogenerated release notes as follows:
Autogenerated release notes as follows:
filter_closure_convert (and new JAX breaking Equinox's experimental stateful operations) by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/232Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.9.1...v0.9.2
These are all pretty self-explanatory!
These are all pretty self-explanatory!
equinox.filter_make_jaxprequinox.filter_vjpequinox.filter_closure_convertequinox.filter_pure_callbackAlso:
equinox.internal.debug_backward_nan(x) will print out the primal and cotangent for x, and if the cotangent has a NaN then the computation is halted.equinox.{is_array, is_array_like, is_inexact_array, is_inexact_array_like} all now recognise NumPy scalars as being array types.equinox.internal.{error_if, branched_error_if} are now compatible with jax.ensure_compile_time_eval.equinox.internal.noinline will now no longer throw an assert error during tracing under certain edge-case conditions. (In particular, when part of the branched of a vmap'd lax.cond with batched predicate.)equinox.tree_pformat now prints out jax.tree_util.Partials, and dataclass types (not instances) correctly.equinox.internal.noinline is now compatible with jax.jit, i.e. a noinline-wrapped function can be passed across a jit API boundary. (Previously equinox.filter_jit was required.)equinox.internal.announce_jaxpr has been renamed to equinox.internal.announce_transform.equinox.internal.{nondifferentiable, nondifferentiable_backward} now take a msg argument for overriding their error messages.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.9.0...v0.9.1
This is a big update. The highlight here is the new equinox.internal namespace, which contains a slew of advanced features.
This is a big update. The highlight here is the new equinox.internal namespace, which contains a slew of advanced features.
These are only "semi public". These are deliberately not in the main documentation, and exist primarily for the benefit of downstream libraries like Diffrax. But you may still have fun playing with them.
equinox.internal.
nondifferentiable: will raise an error at trace-time if you attempt to differentiate it.nondifferentiable_backward: will raise an error at trace-time if you attempt to reverse-mode differentiate it.announce_jaxpr: will call a custom callback whenever it is traced/transformed in a jaxpr. print(<transform stack>) is the default callback.error_if: can raise runtime errors. (Works on CPU; doesn't work on TPU. GPU support may be flaky.)branched_error_if: can raise one of multiple errors, depending on a traced value.nextafter: returns the next floating point number. Unlike jnp.nextafter, it is differentiable.prevbefore: returns the previous floating point number. Is differentiable.noinline: used to mark that a subcomputation should be placed in a separate computation graph, e.g. to avoid compiling the same thing multiple times if it is called repeatedly. Can also be used to iteratively recompile just parts of a computation graph, if the sub/super-graph is the only thing that changes.ω: nice syntax for tree arithmetic. For example (x**ω + y**ω).ω == tree_map(operator.add, x, y). Like tree-math but with nicer syntax.filter_primitive_{def, jvp, transpose, batching, bind}: Define rules for custom primitive that accept arbitrary PyTrees; not just JAX arrays.create_vprim: Autodefines batching rules for higher-order primitives, according to transform(vmap(prim)) == vmap(transform(prim)).str2jax: turns a string into a JAX'able object.unvmap_{any, all, max}: apply reductions whilst ignoring the batch dimension.eqx.{filter_jvp,filter_custom_jvp}eqx.nn.GRUCell will now use its bias term. (Previously it was never adding this.)eqx.filter_eval_shape will no longer promote array-likes to arrays, in either its input or its output.eqx.tree_equal now treats JAX arrays and NumPy arrays as equal.eqx.filter_vmap.Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.8.0...v0.9.0
eqx.{is_array,is_inexact_array} now return True for np.ndarrays rather than False. This is technically a breaking change, hence the new minor version…
The ongoing march of small tweaks progresses.
Main changes this release:
eqx.{is_array,is_inexact_array} now return True for np.ndarrays rather than False. This is technically a breaking change, hence the new minor version bump. Rationale in #202.jaxtyping. Hurrah!Other changes:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.7.1...v0.8.0
Autogenerated release notes as follows:
Autogenerated release notes as follows:
NotImplementedError when computing gradients of stateful models by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/191Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.7.0...v0.7.1
Multiple bugfixes for differentiating through, and serialising, eqx.experimental.BatchNorm.
eqx.experimental.BatchNorm.
eqx.experimental.{BatchNorm,SpectralNorm,StateIndex} then the serialisation format has changed.use_ceil added to all pooling layers.Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.6.0...v0.7.0
Refactor: changed from jax.tree_map to jax.tree_util.tree_map to remove all the deprecation warnings JAX has started giving.
eqx.experimental.{BatchNorm,SpectralNorm,StateIndex} under eqx.tree_{de,}serialise_leaves has been tweaked slightly to avoid an edge-case crash. [This is the reason for the minor version bump to 0.6.0, as this is technically a (very minor) compatibility break.]jax.tree_map to jax.tree_util.tree_map to remove all the deprecation warnings JAX has started giving.eqx.nn.Lambda (for use with eqx.nn.Sequential)eqx.default_{de,}serialise_filter_spec (for use `eqx.tree_{de,}serialise_leaves).BatchNorm crashing under jax.grad.Autogenerated release notes as follows:
Sequential indexable and add tests by @jenkspt in https://github.com/patrick-kidger/equinox/pull/153Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.6...v0.6.0
Autogenerated release notes as follows:
Autogenerated release notes as follows:
{Avg,Max}Pool{1,2,3}D -> {Avg,Max}Pool{1,2,3}d. Removed wrong stride default. by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/135Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.5...v0.5.6
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.4...v0.5.5
Feature: added equinox.filter_eval_shape.
equinox.filter_eval_shape.equinox.filter_pmap(out=...) now supports callable arguments.equinox.filter_vmap(out=...) now properly supports callable arguments. (Previously they were experimental. No part of the API has changed.)equinox.filter_{vmap,pmap}(out=...) not working. (Thanks @jatentaki!)Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.3...v0.5.4
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.2...v0.5.3
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.1...v0.5.2
Adds support for grouped convolutions and transposed convolutions, e.g. equinox.nn.Conv2d(..., groups=...). (Thanks @jatentaki!)
This release:
equinox.nn.GroupNorm.equinox.nn.Conv2d(..., groups=...). (Thanks @jatentaki!)tree_deserialise_leaves having no message.Autogenerated release notes as follows:
filter_{vmap,pmap} by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/84nn.Pool by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/89Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.5.0...v0.5.1
This can be used to create ensembles of models.
This is a big update.
Added filter_vmap.
Added filter_pmap.
Added pooling layers:
eqx.nn.Pooleqx.nn.AvgPool1deqx.nn.AvgPool2deqx.nn.AvgPool3deqx.nn.MaxPool1deqx.nn.MaxPool2deqx.nn.MaxPool3dAdded tree_serialise_leaves and tree_deserialise_leaves.
Added tree_inference, as a convenience for toggling all inference flags through a model.
filter_{jit,grad,value_and_grad} now have an easier-to-use API for specifying which arguments have what behaviour.
(args, kwargs) as a single PyTree, then you can specify a default, args, kwargs separately. In particular this avoids doing messy stuff like filter_spec=((...), {}) when you had no kwargs.args and kwargs against their runtime values. Both the runtime values, and the filter specification, are matched up against the function signature.
e.g. you can do filter_jit(lambda x: x, kwargs=dict(x=True))(1), using a keyword argument at JIT-time and a positional argument at call time.filter_jit(fun) and filter_jit(default=...)(fun) will work.tree_at can now replace subtrees, and not just leaves.
filter, partition now support an is_leaf argument.
filter_jit(filter_grad(fun)) twice will no longer lead to unnecessary recompilation: the second filter_grad(fun) instance will be a PyTree that looks like the first filter_grad(fun) instance, and thus we won't get any recompilation.
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.4.0...v0.5.0
A new minor release as there's a few minor breaking changes:
A new minor release as there's a few minor breaking changes:
MultiheadAttention no longer have a bias by default (#60)equinox.experimental.{get_state,set_state,BatchNorm} now raises RuntimeErrors for many things; this is to match a change in how jax.experimental.host_callback raises errors in jaxlib>=0.3.5. (#63)Besides this, a couple of more exciting (?) things:
equinox.tree_pformat (which is used when printing equinox.Modules) now pretty-prints results much more neatly. (#62)equinox.experimental.{get_state,BatchNorm,SpectralNorm} are now substantially faster when run in inference mode. (#61)Both of which sound pretty minor but both of which were technically really interesting to implement ;)
The pull requests in this release were:
MultiheadAttention by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/60Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.3.2...v0.4.0
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.3.1...v0.3.2
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.3.0...v0.3.1
Removed several old pieces of deprecated functionality: equinox.jitf and so on.
Three main things in this release.
equinox.experimental.BatchNorm. Hurrah, that's nice to have.equinox.experimental.{get_state, set_state, StateIndex}. These are the technology used to update the statistics of BatchNorm without requiring the user to faff around outputting the model themselves. They work by wrapping jax.experimental.host_callback.call to save and load external state on demand. Which is pretty magic, so these should really be used sparingly...equinox.jitf and so on.Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.2.2...v0.3.0
ConvTranspose, ConvTranspose1d, ConvTranspose2d, ConvTranspose3d
Added several new layers:
(Thanks to @andyehrenberg for implementing much of this, and to @lucidrains for reviewing the implementation of attention.)
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.2.1...v0.2.2
Autogenerated release notes as follows:
Autogenerated release notes as follows:
Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.2.0...v0.2.1
First release using GitHub releases. We'll be using this to serve as a changelog.
First release using GitHub releases. We'll be using this to serve as a changelog.
This bumps the minor version 0.1.6 -> 0.2.0 so this is a breaking release. (Admittedly for something pretty minor. See the autogenerated changelog below.)
filter_grad(has_aux=True) returning arguments in the wrong order. by @patrick-kidger in https://github.com/patrick-kidger/equinox/pull/36Full Changelog: https://github.com/patrick-kidger/equinox/compare/v0.1.6...v0.2.0
Nothing published for this version
Nothing published for this version
Your coding agent can read these notes before it upgrades. Set up the MCP server →