NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #2744 most downloaded on PyPI
Flax: A neural network library for JAX designed for flexibility
Last release 10 days ago
24 Sep 2026
Ships fairly regularly
a new release about every 4 weeks
Nearly every release is documented
notes for 59 of the last 60 stable releases
3 versions withdrawn
withdrawn after publishing
7 years old
67 releases · first in 2020
One column per quarter.
verify Flax RoPE using mathematical definition instead of PyTorch by @copybara-service [bot] in #5558
Full Changelog: v0.12.9...v0.12.10
Add vmap reset() regression test for nnx.metrics ( #5483 ) by @chenkuanliao in #5491
Full Changelog: v0.12.8...v0.12.9
Fix intermediates for Gemma example by @vfdev-5 in #5409
Full Changelog: v0.12.7...v0.12.8
Full Changelog: https://github.com/google/flax/compare/v0.12.7...v0.12.8
Add deprecations file by @samanklesaria in #5382
jax.device_put_sharded and jax.device_put_replicated. by @copybara-service[bot] in #5374jax.Array.__contains__ by @copybara-service[bot] in #5404jax.core APIs. by @copybara-service[bot] in #5430P.UNCONSTRAINED in out_shardings in jit. This leads to simplifications in how we lower shardings to stableHLO. by @copybara-service[bot] in #5435Full Changelog: v0.12.6...v0.12.7
jax.device_put_sharded and jax.device_put_replicated. by @copybara-service[bot] in https://github.com/google/flax/pull/5374jax.Array.__contains__ by @copybara-service[bot] in https://github.com/google/flax/pull/5404jax.core APIs. by @copybara-service[bot] in https://github.com/google/flax/pull/5430P.UNCONSTRAINED in out_shardings in jit. This leads to simplifications in how we lower shardings to stableHLO. by @copybara-service[bot] in https://github.com/google/flax/pull/5435Full Changelog: https://github.com/google/flax/compare/v0.12.6...v0.12.7
Annotate wrapped_fn 's self argument in decorator_lift_transform for trace-ability by @copybara-service [bot] in #5298
wrapped_fn's self argument in decorator_lift_transform for trace-ability by @copybara-service[bot] in #5298manual_type: ManualAxisType parameter on ShapedArray to track varying/unreduced/reduced and remove vma: frozenset parameter. by @copybara-service[bot] in #5339Full Changelog: v0.12.5...v0.12.6
wrapped_fn's self argument in decorator_lift_transform for trace-ability by @copybara-service[bot] in https://github.com/google/flax/pull/5298manual_type: ManualAxisType parameter on ShapedArray to track varying/unreduced/reduced and remove vma: frozenset parameter. by @copybara-service[bot] in https://github.com/google/flax/pull/5339Full Changelog: https://github.com/google/flax/compare/v0.12.5...v0.12.6
maintain data/static definition in split / state by @copybara-service [bot] in #5243
jax.config.pmap_shmap_merge. by @copybara-service[bot] in #5289Full Changelog: v0.12.4...v0.12.5
jax.config.pmap_shmap_merge. by @copybara-service[bot] in https://github.com/google/flax/pull/5289Full Changelog: https://github.com/google/flax/compare/v0.12.4...v0.12.5
This release fixes an issue with nnx.List and nnx.Sequential passing all its elements as static when being treated as a pytree.
This release fixes an issue with nnx.List and nnx.Sequential passing all its elements as static when being treated as a pytree.
linen.Partitioned compatible with linen.WeightNorm by @copybara-service[bot] in #5234linen.WeightNorm and linen.with_partitioning compatibility by @copybara-service[bot] in #5236Full Changelog: v0.12.3...v0.12.4
This release fixes an issue with nnx.List and nnx.Sequential passing all its elements as static when being treated as a pytree.
linen.Partitioned compatible with linen.WeightNorm by @copybara-service[bot] in https://github.com/google/flax/pull/5234linen.WeightNorm and linen.with_partitioning compatibility by @copybara-service[bot] in https://github.com/google/flax/pull/5236Full Changelog: https://github.com/google/flax/compare/v0.12.3...v0.12.4
Set numpy<2.4 to fix DeprecationWarning in the CI doctest by @vfdev-5 in #5163
jax.pmap. by @copybara-service[bot] in #5152nnx.pop remove sown attributes. by @samanklesaria in #5133ValueError: Mapped away dimension of inputs passed to vmap should be sharded the same. Got inconsistent axis specs: None vs batch due to split_rngs being replicated. by @copybara-service[bot] in #5189Full Changelog: v0.12.2...v0.12.3
jax.pmap. by @copybara-service[bot] in https://github.com/google/flax/pull/5152nnx.pop remove sown attributes. by @samanklesaria in https://github.com/google/flax/pull/5133ValueError: Mapped away dimension of inputs passed to vmap should be sharded the same. Got inconsistent axis specs: None vs batch due to split_rngs being replicated. by @copybara-service[bot] in https://github.com/google/flax/pull/5189Full Changelog: https://github.com/google/flax/compare/v0.12.2...v0.12.3
[flaxwmt] Small linter fixes. by @copybara-service [bot] in #5012
nnx.PathContains by @thijs-vanweezel in #5094concrete argument to jax.remat by @copybara-service[bot] in #5121Full Changelog: v0.12.1...v0.12.2
nnx.PathContains by @thijs-vanweezel in https://github.com/google/flax/pull/5094concrete argument to jax.remat by @copybara-service[bot] in https://github.com/google/flax/pull/5121Full Changelog: https://github.com/google/flax/compare/v0.12.1...v0.12.2
Variable.value is now deprecated. Consider the following example:
Variable.value is now deprecated. Consider the following example:
import jax.numpy as jnp
import jax
from flax import nnx
my_param = nnx.Param({'a': 0.0})
@nnx.jit
def f(m):
m.value['a'] = 1.0
return mRunning f(my_param) produces Param(value={'a': 0.0}), not Param(value={'a': 1.0}) as before. This is because getting the value parameter new returns a copy of the pytree values (like dict / list). Instead, use the __setitem__ method to update the value:
@nnx.jit
def f(m):
m['a'] = 1.0
return mnnx.Data and nnx.Static annotations are now deprecated. To create nnx.Pytree or nnx.Module dataclasses use the new nnx.dataclass with nnx.data and nnx.static as field descriptors.
# old
@dataclasses.dataclass
class Foo(nnx.Pytree):
a: nnx.Data[int]
b: nnx.Static[str]
# new
@nnx.dataclass
class Foo(nnx.Pytree):
a: int = nnx.data()
b: str = nnx.static()*Norm layer docstrings: axis_index_groups is unused under SPMD jit. by @copybara-service[bot] in #4940ArrayRef creation to the end of Variable creation by @IvyZX in #4980nnx.tabulate by @samanklesaria in #4948jnp.stack instead of np.stack in flax.training.common_utils.stack_forest by @vfdev-5 in #4991jax.experimental.enable_x64 to jax.enable_x64. by @copybara-service[bot] in #5011nnx.set_metadata to in-place change metadata of the state variables of nnx.Modules by @pfackeldey in #5007nnx.Linear in example by @rapsealk in #5023Variable nodes in iter_graph and eliminate recursion. by @copybara-service[bot] in #5035nnx.jit is used with nnx.custom_vjp by @samanklesaria in #5045allow_duplicates argument of nnx.to_arrays. by @dan-zheng in #5083Full Changelog: v0.12.0...v0.12.1
Flax 0.12.0 includes many updates and some important breaking changes to the NNX API.
Flax 0.12.0 includes many updates and some important breaking changes to the NNX API.
nnx.Pytree and therefore nnx.Module are now stricter with regards to attributes that contain Arrays and changing the status of attributes. For example, the code below now fails:
from flax import nnx
import jax
import jax.numpy as jnp
class Foo(nnx.Module):
def __init__(self, use_bias, rngs):
self.layers = [ # ERROR
nnx.Linear(3, 3, rngs=rngs) for _ in range(5)
]
self.bias = None # status = static
if use_bias:
self.bias = nnx.Param(rngs.params.uniform(3,)) # ERRORThis happens for two reasons:
nnx.data. Alternatively, if the container pytree is a list or a dict, you can use nnx.List or nnx.Dict, which additionally allow mixed "data" and "static" elements.nnx.data or nnx.static. Additionally, assigning Arrays or structures with Arrays to static attributes is now an error, as they will not automatically change to data.To fix the above you can just create layers as a List Module which is automatically recognized as data, and be explicit about bias being a data attribute on the first assignment by using nnx.data:
class Foo(nnx.Module):
def __init__(self, use_bias, rngs):
self.layers = nnx.List([ # nnx.data also works but List is recommended
nnx.Linear(3, 3, rngs=rngs) for _ in range(5)
])
self.bias = nnx.data(None)
if use_bias:
self.bias = nnx.Param(rngs.params.uniform(3,))For more information check the Module & Pytree guide.
Variables will now eagerly shard their values when sharding_names metadata is provided. A mesh is required—it can be provided either via passing a mesh metadata attribute or setting the global mesh context via jax.set_mesh. This simplifies the process of sharding a Variable to construction time:
jax.config.update('jax_num_cpu_devices', 8)
mesh = jax.make_mesh((2, 4), ('data', 'model'))
with jax.set_mesh(mesh):
variable = nnx.Param(jnp.ones((16, 32)), sharding_names=(None, 'model'))
print(variable.value.sharding)Eager sharding will also occur when using the nnx.with_partitioning initializer decorator and will automatically extend to the Optimizer. This means that both model and optimizer will be sharded at construction without the need for the somewhat cumbersome nnx.get_partition_spec + jax.lax.with_sharding_constraint + nnx.update pattern:
with jax.set_mesh(mesh):
linear = nnx.Linear(
in_features=16, out_features=16, use_bias=False,
kernel_init=nnx.with_partitioning(
nnx.initializers.lecun_normal(), (None, 'model')
),
rngs=nnx.Rngs(0),
)
optimizer = nnx.Optimizer(linear, optax.adam(1e-3), wrt=nnx.Param)
print(linear.kernel.value.sharding)
print(optimizer.opt_state[0].mu.kernel.value.sharding)For projects that currently rely on other means for sharding, eager sharding can be turned off by passing eager_sharding=False to the Variable constructor, either directly or through initializer decorators like nnx.with_partitioning:
linear = nnx.Linear(
in_features=16, out_features=16, use_bias=False,
kernel_init=nnx.with_partitioning(
nnx.initializers.lecun_normal(), (None, 'model'), eager_sharding=False
),
rngs=nnx.Rngs(0),
)
optimizer = nnx.Optimizer(linear, optax.adam(1e-3), wrt=nnx.Param)
print(linear.kernel.value.sharding)
print(optimizer.opt_state[0].mu.kernel.value.sharding)Eager sharding can also be turned off globally via the flax_always_shard_variable config flag or the FLAX_ALWAYS_SHARD_VARIABLE environment variable:
import flax
flax.config.update('flax_always_shard_variable', False)For more information, check out the Variable eager sharding FLIP.
In-place operators will now raise an error. This is done as part of the push for Variables to be compatible with Tracer semantics:
w = nnx.Variable(jnp.array(0))
w += 1 # ERRORThe fix is to simply operate on the .value property instead:
w.value += 1where argument of jax.numpy reductions. Non-boolean mask inputs have been deprecated for several releases, and will result in an error starting in JAX v0.8.0. by @copybara-service[bot] in #4923flax.config.temp_flip_flag by @IvyZX in #4969Full Changelog: v0.11.2...v0.12.0
nnx.merge now doesn't create a copy of the Variables in the incoming states by default, meaning that the new merged structures holds references to the
nnx.merge now doesn't create a copy of the Variables in the incoming states by default, meaning that the new merged structures holds references to the incoming Variables. This enables new patterns, for example its now possible to create models with the same state but with different runtime behavior:
model = SomeModel(...)
# create eval model
eval_model = nnx.merge(*nnx.split(model)) # same Variables, different structure
eval_model.eval()
model and eval_model share the same Variables and are therefore kept in sync but have different runtime behavior, this avoids having to constantly mutate a single model back and forth between different runtime modes which can be error prone / cause unwanted recompilation.
To keep the old behavior use nnx.merge(..., copy=True).
Full Changelog: https://github.com/google/flax/compare/v0.11.1...v0.11.2
Make Sequential() be identity by @SobhanMP in https://github.com/google/flax/pull/4796
Sequential() be identity by @SobhanMP in https://github.com/google/flax/pull/4796jax.sharding.use_mesh with jax.set_mesh. jax.set_mesh can act as a global setter or a context manager. by @copybara-service[bot] in https://github.com/google/flax/pull/4862Full Changelog: https://github.com/google/flax/compare/v0.11.0...v0.11.1
More on this soon! However, some breaking changes to align with the JAX way of doing things were necessary. Most code should remain intact, however, t…
This version of Flax introduces some changes to improve interop with native JAX and adds support for the new jax.experimental.MutableArray. More on this soon! However, some breaking changes to align with the JAX way of doing things were necessary. Most code should remain intact, however, the following changes deviate from the current behavior:
Rngs in standard layers: all standard layers no longer hold a shared reference to the rngs object given in the constructor, instead they now keep a fork-ed copy of the Rngs or RngStream objects. This impacts Using Rngs in NNX Transforms and Loading Checkpoints with RNGs.model to avoid reference sharing, instead the model must be provided as the first argument to update.split and merge when interacting trivially with raw JAX transforms (state must still be manually propagated if not using MutableArrays, and referential transparency is still an issue). This affects when operating on Pytrees containing NNX Objects with jax.tree.* APIs.Checkout the full NNX 0.10 to NNX 0.11 migration guide.
In the near future we'll share more information about new ways of using NNX with JAX transforms directly by leveraging the new Pytree and MutableArray support. Stay tuned!
.type usage by @vfdev-5 in https://github.com/google/flax/pull/4823.value to [...] in modules_test.py by @lukeyeh in https://github.com/google/flax/pull/4815transforms_test.py from .value to [...] by @lukeyeh in https://github.com/google/flax/pull/4841Full Changelog: https://github.com/google/flax/compare/v0.10.7...v0.11.0
Added identity export from JAX by @jlperla in https://github.com/google/flax/pull/4652
.input_formats and .output_formats in place of .input_layouts and .output_layouts respectively. by @copybara-service in https://github.com/google/flax/pull/4784nnx_basics doc. by @copybara-service in https://github.com/google/flax/pull/4781Full Changelog: https://github.com/google/flax/compare/v0.10.6...v0.10.7
[nnx] remove deprecated APIs by @cgarciae in https://github.com/google/flax/pull/4627
attention_bias parameter to MultiHeadDotProductAttention. by @copybara-service in https://github.com/google/flax/pull/4694attention_bias parameter to MultiHeadDotProductAttention. Add parameter to all overloads to make pytype happy. by @copybara-service in https://github.com/google/flax/pull/4702Full Changelog: https://github.com/google/flax/compare/v0.10.5...v0.10.6
[nnx] fix tabulate by @cgarciae in https://github.com/google/flax/pull/4580
wrappers_test.py to another file. by @copybara-service in https://github.com/google/flax/pull/4581beam_search loop. by @copybara-service in https://github.com/google/flax/pull/4615Full Changelog: https://github.com/google/flax/compare/v0.10.4...v0.10.5
Add deprecation warning to all nnx.State methods by @IvyZX in https://github.com/google/flax/pull/4561
is_initializing API by @copybara-service in https://github.com/google/flax/pull/4550nnx.State methods by @IvyZX in https://github.com/google/flax/pull/4561Full Changelog: https://github.com/google/flax/compare/v0.10.3...v0.10.4
Fix fori_loop and while_loop on multiple modules by @IvyZX in https://github.com/google/flax/pull/4390
async_checkpointer.py reference by @emmanuel-ferdman in https://github.com/google/flax/pull/4385Why Flax NNX documentation by @rajasekharporeddy in https://github.com/google/flax/pull/4425merge docs in graph.py by @8bitmp3 in https://github.com/google/flax/pull/4411nnx.bridge.variables.nnx_attrs_to_linen_vars take nnx.VariableState as argument. by @copybara-service in https://github.com/google/flax/pull/4473Param(None) lines from NNX by @IvyZX in https://github.com/google/flax/pull/4504nnx.Module.perturb by @IvyZX in https://github.com/google/flax/pull/4515Full Changelog: https://github.com/google/flax/compare/v0.10.2...v0.10.3
Add nnx.fori_loop by @IvyZX in https://github.com/google/flax/pull/4353
nnx.fori_loop by @IvyZX in https://github.com/google/flax/pull/4353nn.jit under nn.scan. by @copybara-service in https://github.com/google/flax/pull/4359flax.nnx.eval_shape docstring by @8bitmp3 in https://github.com/google/flax/pull/4374flax.nnx.remat docstring by @8bitmp3 in https://github.com/google/flax/pull/4373Full Changelog: https://github.com/google/flax/compare/v0.10.1...v0.10.2
Add Flax NNX GraphDef docstring by @8bitmp3 in https://github.com/google/flax/pull/4302
nnx.while_loop and nnx.switch by @IvyZX in https://github.com/google/flax/pull/4343Full Changelog: https://github.com/google/flax/compare/v0.10.0...v0.10.1
Forward all arguments when using nnx.transforms.deprecated.scan as a decorator. by @copybara-service in https://github.com/google/flax/pull/4208
nnx.bridge by @IvyZX in https://github.com/google/flax/pull/4145nnx.bridge by @IvyZX in https://github.com/google/flax/pull/4171nnx.Variable by @IvyZX in https://github.com/google/flax/pull/4189impost to import by @Vilin97 in https://github.com/google/flax/pull/4256docs_nnx doctest by @IvyZX in https://github.com/google/flax/pull/4285Full Changelog: https://github.com/google/flax/compare/v0.9.0...v0.10.0
Ignore Orbax warning in deprecated flax.training.checkpoints.py to unbreak head doctest by @IvyZX in https://github.com/google/flax/pull/4092
nnx_basics doc by @rajasekharporeddy in https://github.com/google/flax/pull/4047fp32_max_grad by @kaixih in https://github.com/google/flax/pull/3984nnx.compat to nnx.bridge by @IvyZX in https://github.com/google/flax/pull/4066flax.training.checkpoints.py to unbreak head doctest by @IvyZX in https://github.com/google/flax/pull/4092LinenToNNX (prev LinenWrapper) by @IvyZX in https://github.com/google/flax/pull/4081einsum_dot_general and einsum argument by specifying them explicitly in the caller. by @copybara-service in https://github.com/google/flax/pull/4115nnx.bridge by @IvyZX in https://github.com/google/flax/pull/4126_ParentType annotation by @dcharatan in https://github.com/google/flax/pull/4120Full Changelog: https://github.com/google/flax/compare/v0.8.5...v0.9.0
fix deprecation warning by @chiamp in https://github.com/google/flax/pull/3981
nnx.module docstrings by @chiamp in https://github.com/google/flax/pull/3966nnx.Conv and nnx.ConvTranspose by @chiamp in https://github.com/google/flax/pull/3974nnx.graph docstrings by @chiamp in https://github.com/google/flax/pull/3958pmap and Pmap. static_broadcasted_argnums, donate_argnums, and global_arg_shapes are not yet supported. by @copybara-service in https://github.com/google/flax/pull/3978rnglib docstring by @chiamp in https://github.com/google/flax/pull/3980nnx.training by @chiamp in https://github.com/google/flax/pull/3975nnx.variables docstrings by @chiamp in https://github.com/google/flax/pull/3986wrt option to nnx.Optimizer by @chiamp in https://github.com/google/flax/pull/3983nnx.graph.iter_children by @chiamp in https://github.com/google/flax/pull/3991force_fp32_for_softmax arg in MultiHeadDotProductAttention useful. by @copybara-service in https://github.com/google/flax/pull/4029Full Changelog: https://github.com/google/flax/compare/v0.8.4...v0.8.5
fixed jnp.clip deprecation by @chiamp in https://github.com/google/flax/pull/3905
jnp.clip deprecation by @chiamp in https://github.com/google/flax/pull/3905codediff and added testing for first tab by @chiamp in https://github.com/google/flax/pull/3847jax.sharding.PartitionSpec.UNCONSTRAINED in logical specification by @copybara-service in https://github.com/google/flax/pull/3902Module.scope.path by @chiamp in https://github.com/google/flax/pull/3913jax.tree_* functions with jax.tree.* by @copybara-service in https://github.com/google/flax/pull/3926nnx.ConvTranspose by @chiamp in https://github.com/google/flax/pull/3934nnx.Sequential by @chiamp in https://github.com/google/flax/pull/3935LoRA and LoRALinear to NNX by @IvyZX in https://github.com/google/flax/pull/3929Full Changelog: https://github.com/google/flax/compare/v0.8.3...v0.8.4
added tree_map deprecation warning filter by @chiamp in https://github.com/google/flax/pull/3828
nnx.Pytree by @chiamp in https://github.com/google/flax/pull/3743TrainState's step possibly jax.Array. This makes replicate valid for type checking. by @copybara-service in https://github.com/google/flax/pull/3763reset_gate test by @chiamp in https://github.com/google/flax/pull/3773kw_only struct.dataclass test by @chiamp in https://github.com/google/flax/pull/3651PyTreeNode to take dataclass kwargs by @chiamp in https://github.com/google/flax/pull/3785nnx.GraphNode by @chiamp in https://github.com/google/flax/pull/3796nnx.training by @chiamp in https://github.com/google/flax/pull/3782grads kwarg for Optimizer.update by @chiamp in https://github.com/google/flax/pull/3818tree_map deprecation warning filter by @chiamp in https://github.com/google/flax/pull/3828tree_map by @chiamp in https://github.com/google/flax/pull/3823robots.txt by @chiamp in https://github.com/google/flax/pull/3886Full Changelog: https://github.com/google/flax/compare/v0.8.2...v0.8.3
Replace deprecated API jax.tree_map by @copybara-service in https://github.com/google/flax/pull/3715
jax.tree_map by @copybara-service in https://github.com/google/flax/pull/3715jax.tree_util.tree_map instead of deprecated jax.tree_map. by @copybara-service in https://github.com/google/flax/pull/3714strides and kernel_dilation to nn.ConvTranspose by @IvyZX in https://github.com/google/flax/pull/3731Full Changelog: https://github.com/google/flax/compare/v0.8.1...v0.8.2
fixed rng guide outputs by @chiamp in #3685
enforce mask kwarg in norm layers by @chiamp in #3663
added kwargs to self.param and self.variable by @chiamp in #3675
added nnx normalization tests by @chiamp in #3689
added NNX init_cache docstring example by @chiamp in #3688
added nnx attention equivalence test by @chiamp in #3687
Fix bug that assumed frozen-dict keys were strings. by @copybara-service in #3692
added nnx rmsnorm by @chiamp in #3691
updated nnx compute_stats by @chiamp in #3693
fixed intercept_methods docstring by @chiamp in #3694
[nnx] Add Sphinx Docs by @cgarciae in #3678
Fix pointless docstring example of nn.checkpoint / nn.remat. by @levskaya in #3703
added default params rng to .apply by @chiamp in #3698
[nnx] add partial_init by @cgarciae in #3674
make make_rng default to 'params' by @chiamp in #3699
Add SimpleCell. by @carlosgmartin in #3697
fix Module.module_paths docstring by @cgarciae in #3709
Guarantee the latest JAX version on CI by @cgarciae in #3705
Replace deprecated API jax.tree.map by @copybara-service in #3715
Use jax.tree_util.tree_map instead of deprecated jax.tree.map . by @copybara-service in #3714
[nnx] simplify readme by @cgarciae in #3707
[nnx] add demo.ipynb by @cgarciae in #3680
Fix Tabulate's compute_flops by @cgarciae in #3721
[nnx] simplify TraceState by @cgarciae in #3724
Add broadcast of strides and kernel_dilation to nn.ConvTranspose by @IvyZX in #3731
[nnx] Fix State. sub by @cgarciae in #3704
[nnx] always fold_in on fork + new ForkedKeys return type by @cgarciae in #3722
[nnx] explicit Variables by @cgarciae in #3720
Improves fingerprint definition for Modules in nn.jit. by @copybara-service in #3736
Flax: avoid key reuse in tests by @copybara-service in #3740
added Einsum layer by @chiamp in #3710
nn.jit: automatic fingerprint definition for dataclass attributes by @cgarciae in #3737
[NVIDIA] Use custom grad accumulation for FP8 params by @kaixih in #3623
removed nnx dataclass by @chiamp in #3742
[nnx] cleanup graph_utils by @cgarciae in #3728
Fix doctest and unbreak head by @IvyZX in #3753
[nnx] add pytree support by @cgarciae in #3732
fixed intercept_methods docstring by @chiamp in #3752
Add ConvLSTMCell to docs. by @carlosgmartin in #3712
[nnx] remove flagslib by @cgarciae in #3733
Fix tests after applying JAX key-reuse checker. See: by @copybara-service in #3748
bump version number to 0.8.1 by @chiamp in https://github.com/google/flax/pull/3649
Full Changelog: https://github.com/google/flax/compare/v0.8.0...v0.8.1
make_rng.InstanceNorm and renamed channel_axes to feature_axes.Module.module_paths and doc.Sequential.__call__ compact.nn.compact_name_scope v3.flax.struct.dataclass.jax.tree_util.tree_map with mapping over leafs.exempt a jax.config deprecation warning by @levskaya in https://github.com/google/flax/pull/3465
PReLU Test by @Micky774 in https://github.com/google/flax/pull/3498Embed layer by @Micky774 in https://github.com/google/flax/pull/3513Conv layer by @Micky774 in https://github.com/google/flax/pull/3511Linear/Dense layer by @Micky774 in https://github.com/google/flax/pull/3509Conv NNX/Linen consistency test by @Micky774 in https://github.com/google/flax/pull/3526_hashable_filter does not convert strings to a tuple of letters by @copybara-service in https://github.com/google/flax/pull/3533return_weights to sow_weights for attention layer by @chiamp in https://github.com/google/flax/pull/3550Full Changelog: https://github.com/google/flax/compare/v0.7.5...v0.8.0
nn.compact_name_scope decorator that enables methods to act as compact name scopes as with regular Haiku methods. This makes porting Haiku code easier.BatchApply class.sow_weights option in attention layer.MultiHeadAttention alias.nn.jit.normalize activation function, in favor of standardize.GeGLU activation function.Enum support for tabulate function.nn.grad function.fix DeprecationWarnings by @chiamp in https://github.com/google/flax/pull/3352
find methods and magic methods for Cursor API by @chiamp in https://github.com/google/flax/pull/3306MultiHeadDotProductAttention by @chiamp in https://github.com/google/flax/pull/3384has_improved field to EarlyStopping by @chiamp in https://github.com/google/flax/pull/3385Full Changelog: https://github.com/google/flax/compare/v0.7.4...v0.7.5
linen.Module.tabulate and summary.tabulate (in new flops and vjp_flops table columns). Pass compute_flops=True and/or compute_vjp_flops=True to include these columns.MultiHeadDotProductAttention's call method signature, by adding
inputs_k and inputs_v args and switching inputs_kv, mask and determistic
to keyword arguments. See more details in #3389.jax.random.PRNGKey to jax.random.key.
(See JEP 9263 for details).
If you notice dispatch performance regressions after this change, be sure
you update jax to version 0.4.16 or newer.has_improved field to EarlyStopping and changed the return signature of
EarlyStopping.update from returning a tuple to returning just the updated class.
See more details in #3385Added python version constraint >=3.9.
Added python version constraint >=3.9.
Full Changelog: https://github.com/google/flax/compare/v0.7.3...v0.7.4
New features:
Bug fixes:
Fix Documentation Typo by @peterdavidfagan in https://github.com/google/flax/pull/3265
Full Changelog: https://github.com/google/flax/compare/v0.7.2...v0.7.3
make flax.core.copy add_or_replace optional by @PhilipVinc in https://github.com/google/flax/pull/3241
add_or_replace optional by @PhilipVinc in https://github.com/google/flax/pull/3241Full Changelog: https://github.com/google/flax/compare/v0.7.1...v0.7.2
New features:
flax.core.copy add_or_replace optionaluse_fast_variance option to GroupNorm and BatchNorm to allow disabling it.Bug fixes:
field_specifiers instead of field_descriptors in @dataclass_transform.nn.Module typing.jax.experimental.pjit.with_sharding_constraint with jax.lax.with_sharding_constraint.added dict migration guide to index by @chiamp in https://github.com/google/flax/pull/3188
LayerNorm. by @copybara-service in https://github.com/google/flax/pull/3194compute_stats. by @copybara-service in https://github.com/google/flax/pull/3205Full Changelog: https://github.com/google/flax/compare/v0.7.0...v0.7.1
Breaking changes:
New features:
target arrays when they port the target shardings.Bug fixes:
orbax.checkpoint which is a better import pattern.orbax.checkpoint as ocp to avoid the verbosity of using 'orbax.checkpoint` every time.LayerNorm.stack instead of concatenate in compute_stats, to handle scalar stats case.compute_stats.delete long-unused dotgetter utility by @copybara-service in https://github.com/google/flax/pull/3156
restore_with_serialized_types in preparation for an upcoming change. by @copybara-service in https://github.com/google/flax/pull/3165Full Changelog: https://github.com/google/flax/compare/v0.6.11...v0.7.0
remove unreferenced docs by @chiamp in https://github.com/google/flax/pull/3085
pjit guide to jax.jit by @IvyZX in https://github.com/google/flax/pull/3016Full Changelog: https://github.com/google/flax/compare/v0.6.10...v0.6.11
Add missing Bidirectional api reference by @cgarciae in https://github.com/google/flax/pull/2999
urllib3 error by @chiamp in https://github.com/google/flax/pull/3081Full Changelog: https://github.com/google/flax/compare/v0.6.9...v0.6.10
Follow-up to #2985 (adding config.prefetch) by @copybara-service in https://github.com/google/flax/pull/3004
config.prefetch) by @copybara-service in https://github.com/google/flax/pull/3004get_partition_spec on replicated array by @IvyZX in https://github.com/google/flax/pull/3013pretty_repr utility fn by @chiamp in https://github.com/google/flax/pull/3022Full Changelog: https://github.com/google/flax/compare/v0.6.8...v0.6.9
orbax-checkpoint package instead of orbax.project.toml.jax.sharding, and with more warnings.The automatic checkpoint migration was temporarily rolled back due to legacy compatibility issues.
flax.config.update('flax_use_orbax_checkpointing', True) to your project to avoid being impacted by the automatic migration process.register_keypaths.register_keypaths by @copybara-service in https://github.com/google/flax/pull/2963Full Changelog: https://github.com/google/flax/compare/v0.6.7...v0.6.8
Fix lm1b sampler by @cgarciae in https://github.com/google/flax/pull/2842
Full Changelog: https://github.com/google/flax/compare/v0.6.6...v0.6.7
flax.config.update('flax_use_orbax_checkpointing', False) to temporarily disable this migration, but note that Flax legacy checkpointing will be removed 3 months from Mar 10, 2023.FrozenDict to regular dict: utility functions now work on both.FrozenDict to JAX pytree keypath API.ModuleEnhance checkpointing docs with Orbax information by @cpgaffney1 in https://github.com/google/flax/pull/2879
Full Changelog: https://github.com/google/flax/compare/v0.6.5...v0.6.6
Internal change by @copybara-service in https://github.com/google/flax/pull/2826
zeros and ones initializers with zeros_init and ones_init builder functions by @chiamp in https://github.com/google/flax/pull/2815Full Changelog: https://github.com/google/flax/compare/v0.6.4...v0.6.5
Module.lazy_init to avoid compute during Module initialization.Add deprecation warnings and docstring pointers to flax.training.lr_schedule by @IvyZX in https://github.com/google/flax/pull/2702
Full Changelog: https://github.com/google/flax/compare/v0.6.3...v0.6.4
New features:
jax.pjit: https://flax.readthedocs.io/en/latest/guides/flax_on_pjit.htmlzeros_init and ones_init initializers.flax.config flags.Bug fixes:
dataclass.fields from __repr__.Add gfile api shim to remove tensorflow dependency for basic IO. by @chiamp in https://github.com/google/flax/pull/2586
Full Changelog: https://github.com/google/flax/compare/v0.6.2...v0.6.3
New features:
Bug fixes:
Refactor out dataclass transform that allows parent and name to be moved to the end of the argument list into a more general "kw_only_dataclasses" mod
parent and name to be moved to the end of the argument list into a more general "kw_only_dataclasses" module. by @copybara-service in https://github.com/google/flax/pull/2468is_fully_replicated and is_fully_addressble a property rather than a method. by @copybara-service in https://github.com/google/flax/pull/2516gfile.remove for files because it doesn't work on GCS files. by @IvyZX in https://github.com/google/flax/pull/2518is_initializing for init-detection. by @yotarok in https://github.com/google/flax/pull/2486launch_gce.sh with imagenet example. by @andsteing in https://github.com/google/flax/pull/2598Full Changelog: https://github.com/google/flax/compare/v0.6.1...v0.6.2
New features:
gfile.remove for files because it doesn't work on GCS files.Bug fixes:
pmap.ignore tf deprecation warning. by @copybara-service in https://github.com/google/flax/pull/2466
rich dependency version by @yklcs in https://github.com/google/flax/pull/2407Full Changelog: https://github.com/google/flax/compare/v0.6.0...v0.6.1
Add on_commit_callback to put the responsibility of renaming the directories on the users of the serialization library. This will also fix the GCS ato
on_commit_callback to put the responsibility of renaming the directories on the users of the serialization library. This will also fix the GCS atomic rename issue where the users can write a success file when the commit is successful and check the existence of that file before deserialization. by @copybara-service in https://github.com/google/flax/pull/2328step in training.checkpoints by @Chuxiaof in https://github.com/google/flax/pull/2343dynamic_scale.py from optim/ to training/. by @copybara-service in https://github.com/google/flax/pull/2375auto_flush to flax.metrics.tensorboard.SummaryWriter by @copybara-service in https://github.com/google/flax/pull/2376flax.optim.dynamic_scale. by @copybara-service in https://github.com/google/flax/pull/2314Full Changelog: https://github.com/google/flax/compare/v0.5.3...v0.6.0
flax.optim package.flax.optim.dynamic_scale to flax.training.dynamic_scale.jax.named_scope for all profile naming, cut some pointless
stack traces out.Update reference to tree_map to avoid deprecation warning. by @copybara-service in https://github.com/google/flax/pull/2298
flax.linen.Module.bind by @nalzok in https://github.com/google/flax/pull/2269Full Changelog: https://github.com/google/flax/compare/v0.5.2...v0.5.3
New features:
nn.switch as a lifted version of jax.lax.switch.jax.experimental.GlobalDeviceArray, a useful array type for multiprocess/multihost computing.save_checkpoints() on single-process scenario.Bug fixes:
Flax Basics docs: Add missing @jax.jit to mse by @rsokl in https://github.com/google/flax/pull/2181
@jax.jit to mse by @rsokl in https://github.com/google/flax/pull/2181Full Changelog: https://github.com/google/flax/compare/v0.5.1...v0.5.2
Adds flax import to summary.py by @marcvanzee in https://github.com/google/flax/pull/2138
Full Changelog: https://github.com/google/flax/compare/v0.5.0...v0.5.1
New features:
nn.tabulate and Module.tabulate to generate rich representations of the network structure.Added flax.jax_utils.ad_shard_unpad() by @lucasb-eyer
New features:
flax.jax_utils.ad_shard_unpad() by @lucasb-eyerBug fixes:
Breaking changes:
Use tree_map instead of deprecated tree_multimap by @jheek in https://github.com/google/flax/pull/2024
optim.Adam(weight_decay) parameter. by @copybara-service in https://github.com/google/flax/pull/2016Full Changelog: https://github.com/google/flax/compare/v0.4.1...v0.4.2
New features:
nn.cond.jax.tree_multimap (deprecated) with jax.tree.map.Bug fixes:
Added Optax update guide and deprecated flax.optim.
flax.linen.ConvLocal.flax.optim.sep argument to flax.traverse_util.flatten_dict().flax.linen.combinators.Removes deprecated API from RTD by @marcvanzee in https://github.com/google/flax/pull/1824
unroll to jax_utils.scan_in_dim by @ptigwe in https://github.com/google/flax/pull/1691rng arguments from Dropout's __call__. by @copybara-service in https://github.com/google/flax/pull/1689scenic link to examples/README.md by @copybara-service in https://github.com/google/flax/pull/1809Full Changelog: https://github.com/google/flax/compare/v0.3.6...v0.4.0
Breaking changes:
New features:
flax.linen.custom_vjp for custom derivatives inside a Module.param_dtype attribute to standard Linen Modules for specifying parameter dtypes.Move flax.nn to flax.deprecated.nn.
Breaking changes:
flax.nn to flax.deprecated.nn.New features:
flax.linen.checkpointflax.linen.map_variables.You can no longer pass an int as the kernel_size for a flax.linen.Conv. Instead a type error is raised stating that a tuple/list should be provided. S
Breaking changes:
flax.linen.Conv. Instead a type error is raised stating that a tuple/list should be provided. Stride and dilation arguments do support broadcasting a single int value now because this is not ambiguous when the kernel rank is known.--jax_numpy_rank_promotion=raise.Bugfixes:
When calling init the 'intermediates' collection is no longer mutable. Therefore, intermediates will no longer be returned from initialization by defa
Possibly breaking changes:
init the 'intermediates' collection is no longer mutable.
Therefore, intermediates will no longer be returned from initialization by default.deterministic argument in MultiHeadDotProductAttention.Other changes:
examples/sst2.
that uses a bidirectional LSTM (BiLSTM) to encode the input text.flax.training.train_state to simplify using Optax optimizers.mutable argument is now available on Module.init and Module.init_with_outputsdot_product_attention_weights, allowing access to attention weights.BatchNorm instances will behave correctly during init when called multiple times.contributing.md.lift.jit,
fixing cache misses.linen.Module for deep inheritance chains.MultiOptimizer use apply_gradient instead of apply_param_gradient.Bug Fix: Disallow modifying attributes in Modules after they are initialized.
Possible breaking changes:
Other changes:
Module.bind for binding variables and RNGs to an interactive Module.nn.apply and nn.init for transforming arbitrary functions that take a linen.Module as their first argument.save_checkpoint.is_leaf argument in traverse_util.flatten_dictFixed an overly aggressive flax.nn deprecation warning in v0.3.1 that fired whenever you imported flax.
Fixed an overly aggressive flax.nn deprecation warning in v0.3.1 that fired whenever you imported flax.
flax.nn deprecation message no longer appears if you import flax directly.
NOTE: You must now explicitly import flax.nn if you want to use the old
pre-Linen flax.nn.Module.
Many improvements to Linen, and the old flax.nn is officially reprecated!
Many improvements to Linen, and the old flax.nn is officially reprecated!
Notably, there's a clean API for extracting intermediates from modules
defined using @nn.compact, a more ergonomic API for using Batch Norm and Dropout in modules
defined using setup, support for MultiOptimizer with Linen, and multiple safety, performance
and error message improvements.
See the CHANGELOG for more details
Many improvements to Linen, and the old flax.nn is officially deprecated!
Notably, there's a clean API for extracting intermediates from modules
defined using @nn.compact, a more ergonomic API for using Batch Norm and Dropout in modules
defined using setup, support for MultiOptimizer with Linen, and multiple safety, performance
and error message improvements.
Possible breaking changes:
Module instances are now frozen after setup has been called.
Previously mutations after setup could be dropped silently. Now the stateless requirement
is enforced by raising a TypeError in __setattr__ after setup.flax.core.apply and Module.apply. Now it returns a tuple
containing the output and a frozen empty
collection when mutable is specified as an empty list.broadcast_dims is now an attribute to Dropout instead of a __call__
argument.use_running_average and deterministic no longer have a default. They
should be passed explicitlyScope.variable mutability check, before a variable could only be
initialized if the 'params' collection was mutable.Other Improvements:
lm1b language modeling exampleflax.nn classes and methods.__call__. See design notesow method to Module and capture_intermediates argument to Module.apply.
See howto for usage patterns.variable argument to Module.apply to be a FrozenDictModelParamTraversal
As a result MultiOptimizer can be used properly with linen modules.is_mutable() method to Variable and is_mutable_collection() to flax.linen.Module.axis_name arg to flax.linen.vmapflax.linen.scanYour coding agent can read these notes before it upgrades. Set up the MCP server →