flax
Flax is a neural network library for JAX that is designed for flexibility.
File Explorer
Download Latest Version (.zip)- dependabot.yml
- zizmor.yml
- get_repo_metrics.py
- issue_activity_since_date.gql
- pr_data_query.gql
- README.md
- requirements.txt
- bug_report.md
- flax_publish.yml
- flax_test.yml
- flaxlib_publish.yml
- jax_nightly.yml
- pull_request_template.md
- __init__.py
- gemma.py
- imagenet.py
- lm1b.py
- mnist.py
- nlp_seq.py
- ogbg_molpcba.py
- ppo.py
- README.md
- requirements.txt
- run_all_benchmarks.sh
- seq2seq.py
- sst2.py
- tracing_benchmark.py
- vae.py
- wmt.py
- nnx_graph_overhead.py
- nnx_mlpmixer_training.py
- nnx_simple_training.py
- nnx_state_traversal.py
- README.md
- codediff.py
- codediff_test.py
- flax_module.py
- flax_theme.css
- flax_module.rst
- activation_functions.rst
- decorators.rst
- index.rst
- init_apply.rst
- initializers.rst
- inspection.rst
- layers.rst
- module.rst
- profiling.rst
- spmd.rst
- transformations.rst
- variable.rst
- flax.core.frozen_dict.rst
- flax.cursor.rst
- flax.errors.rst
- flax.jax_utils.rst
- flax.serialization.rst
- flax.struct.rst
- flax.traceback_util.rst
- flax.training.rst
- index.rst
- index.rst
- lift.md
- module_lifecycle.rst
- community_examples.rst
- core_examples.rst
- google_research_examples.rst
- index.rst
- repositories_that_use_flax.rst
- 0000-template.md
- 1009-optimizer-api.md
- 1777-default-dtype.md
- 2396-rnn.md
- 2434-general-metadata.md
- 2974-kw-only-dataclasses.md
- 3099-rnnbase-refactor.md
- 4105-jax-style-nnx-transforms.md
- README.md
- convert_pytorch_to_flax.rst
- haiku_migration_guide.rst
- index.rst
- linen_upgrade_guide.rst
- optax_update_guide.rst
- orbax_upgrade_guide.rst
- regular_dict_upgrade_guide.rst
- rnncell_upgrade_guide.rst
- full_eval.rst
- index.rst
- loading_datasets.ipynb
- loading_datasets.md
- arguments.md
- flax_basics.ipynb
- flax_basics.md
- index.rst
- rng_guide.ipynb
- rng_guide.md
- setup_or_nncompact.rst
- state_params.rst
- extracting_intermediates.rst
- index.rst
- model_surgery.ipynb
- model_surgery.md
- ensembling.rst
- flax_on_pjit.ipynb
- flax_on_pjit.md
- index.rst
- fp8_basics.ipynb
- fp8_basics.md
- index.rst
- batch_norm.rst
- dropout.rst
- index.rst
- lr_schedule.rst
- transfer_learning.ipynb
- transfer_learning.md
- use_checkpointing.ipynb
- use_checkpointing.md
- flax_sharp_bits.ipynb
- flax_sharp_bits.md
- index.rst
- .gitignore
- .readthedocs.yaml
- conf.py
- conf_sphinx_patch.py
- faq.rst
- flax.png
- glossary.rst
- index.rst
- linen_intro.ipynb
- linen_intro.md
- Makefile
- quick_start.ipynb
- quick_start.md
- README.md
- robots.txt
- codediff.py
- codediff_test.py
- flax_module.py
- flax_theme.css
- flax_module.rst
- activations.rst
- attention.rst
- dtypes.rst
- index.rst
- initializers.rst
- linear.rst
- lora.rst
- normalization.rst
- pooling.rst
- recurrent.rst
- stochastic.rst
- ema.rst
- index.rst
- metrics.rst
- optimizer.rst
- bridge.rst
- compat.rst
- filterlib.rst
- graph.rst
- helpers.rst
- index.rst
- module.rst
- object.rst
- rnglib.rst
- spmd.rst
- state.rst
- summary.rst
- transforms.rst
- variables.rst
- visualization.rst
- flax.config.rst
- flax.core.frozen_dict.rst
- flax.struct.rst
- flax.training.rst
- flax.traverse_util.rst
- index.rst
- digits_diffusion_model.ipynb
- digits_diffusion_model.md
- gemma.ipynb
- gemma.md
- image_segmentation.ipynb
- image_segmentation.md
- index.rst
- machine_translation.ipynb
- machine_translation.md
- minigpt.ipynb
- minigpt.md
- object_detection_detr.ipynb
- object_detection_detr.md
- vit_training.ipynb
- vit_training.md
- 0000-template.md
- 1009-optimizer-api.md
- 1777-default-dtype.md
- 2396-rnn.md
- 2434-general-metadata.md
- 2974-kw-only-dataclasses.md
- 3099-rnnbase-refactor.md
- 4105-jax-style-nnx-transforms.md
- 4844-var-eager-sharding.md
- 5310-tree-mode-nnx.md
- README.md
- performance-graph.png
- stateful-transforms.png
- blog.md
- bridge_guide.ipynb
- bridge_guide.md
- checkpointing.ipynb
- checkpointing.md
- data_loaders.ipynb
- data_loaders.md
- demo.ipynb
- demo.md
- extracting_intermediates.ipynb
- extracting_intermediates.md
- filters_guide.ipynb
- filters_guide.md
- flax_gspmd.ipynb
- flax_gspmd.md
- index.rst
- jax_and_nnx_transforms.rst
- optimization_cookbook.ipynb
- optimization_cookbook.md
- performance.ipynb
- performance.md
- pytree.ipynb
- pytree.md
- randomness.ipynb
- randomness.md
- surgery.ipynb
- surgery.md
- tiny_nnx.ipynb
- transforms.ipynb
- transforms.md
- view.ipynb
- view.md
- hijax.ipynb
- hijax.md
- index.rst
- haiku_to_flax.rst
- index.rst
- linen_to_nnx.rst
- nnx_010_to_nnx_011.rst
- pytorch_to_jax_flax.rst
- .gitignore
- .readthedocs.yaml
- conf.py
- conf_sphinx_patch.py
- contributing.md
- faq.rst
- flax.png
- guides_advanced.rst
- guides_basic.rst
- index.rst
- key_concepts.ipynb
- key_concepts.md
- Makefile
- mnist_tutorial.ipynb
- mnist_tutorial.md
- nnx_basics.ipynb
- nnx_basics.md
- nnx_glossary.rst
- philosophy.md
- README.md
- robots.txt
- tensorboard_screenshot.png
- why.rst
- launch_gce.py
- README.md
- startup_script.sh
- gemma3_1b_grain.py
- gemma3_1b_tf.py
- default.py
- gemma3_270m.py
- gemma3_270m_sow.py
- gemma3_4b.py
- small.py
- tiny.py
- helpers.py
- helpers_test.py
- input_pipeline.py
- input_pipeline_grain.py
- input_pipeline_test.py
- input_pipeline_tf.py
- layers.py
- layers_test.py
- main.py
- modules.py
- modules_test.py
- params.py
- positional_embeddings.py
- positional_embeddings_test.py
- README.md
- requirements.txt
- sampler.py
- sampler_test.py
- sow_lib.py
- tokenizer.py
- train.py
- train_cfg.py
- transformer.py
- transformer_cfg.py
- transformer_test.py
- default.py
- fake_data_benchmark.py
- tpu.py
- v100_x8.py
- v100_x8_mixed_precision.py
- imagenet.ipynb
- imagenet_benchmark.py
- imagenet_fake_data_benchmark.py
- input_pipeline.py
- main.py
- models.py
- models_test.py
- README.md
- requirements.txt
- train.py
- train_test.py
- attention_simple.py
- autoencoder.py
- dense.py
- linear_regression.py
- mlp_explicit.py
- mlp_inline.py
- mlp_lazy.py
- default.py
- input_pipeline.py
- input_pipeline_test.py
- main.py
- models.py
- README.md
- requirements.txt
- temperature_sampler.py
- temperature_sampler_test.py
- tokenizer.py
- train.py
- train_test.py
- utils.py
- default.py
- main.py
- mnist.ipynb
- mnist_benchmark.py
- README.md
- requirements.txt
- train.py
- train_test.py
- default.py
- input_pipeline.py
- input_pipeline_test.py
- main.py
- models.py
- README.md
- requirements.txt
- train.py
- 01_functional_api.py
- 02_lifted_transforms.py
- 03_train_state.py
- 04_data_parallel_with_jit.py
- 05_vae.py
- 06_scan_over_layers.py
- 07_array_leaves.py
- 08_save_load_checkpoints.py
- 09_parameter_surgery.py
- 10_fsdp_and_optimizer.py
- hijax_basic.py
- hijax_demo.py
- requirements.txt
- default.py
- default_graph_net.py
- hparam_sweep.py
- test.py
- input_pipeline.py
- input_pipeline_test.py
- main.py
- models.py
- models_test.py
- ogbg_molpcba.ipynb
- ogbg_molpcba_benchmark.py
- README.md
- requirements.txt
- train.py
- train_test.py
- default.py
- agent.py
- env_utils.py
- models.py
- ppo_lib.py
- ppo_lib_test.py
- ppo_main.py
- README.md
- requirements.txt
- seed_rl_atari_preprocessing.py
- test_episodes.py
- default.py
- input_pipeline.py
- main.py
- models.py
- README.md
- requirements.txt
- seq2seq.ipynb
- train.py
- train_test.py
- default.py
- build_vocabulary.py
- input_pipeline.py
- input_pipeline_test.py
- main.py
- models.py
- models_test.py
- README.md
- requirements.txt
- sst2.ipynb
- train.py
- train_test.py
- vocab.txt
- vocabulary.py
- default.py
- .gitignore
- input_pipeline.py
- main.py
- models.py
- README.md
- reconstruction.png
- requirements.txt
- sample.png
- train.py
- utils.py
- default.py
- bleu.py
- decode.py
- input_pipeline.py
- input_pipeline_test.py
- main.py
- models.py
- README.md
- requirements.txt
- tokenizer.py
- train.py
- train_test.py
- __init__.py
- README.md
- __init__.py
- attention.py
- linear.py
- normalization.py
- stochastic.py
- __init__.py
- axes_scan.py
- flax_functional_engine.ipynb
- frozen_dict.py
- lift.py
- meta.py
- partial_eval.py
- scope.py
- spmd.py
- tracers.py
- variables.py
- __init__.py
- nnx.py
- layers_with_named_axes.py
- __init__.py
- activation.py
- attention.py
- batch_apply.py
- combinators.py
- dtypes.py
- fp8_ops.py
- initializers.py
- kw_only_dataclasses.py
- linear.py
- module.py
- normalization.py
- partitioning.py
- pooling.py
- README.md
- recurrent.py
- spmd.py
- stochastic.py
- summary.py
- transforms.py
- __init__.py
- tensorboard.py
- __init__.py
- interop.py
- module.py
- variables.py
- wrappers.py
- __init__.py
- activations.py
- attention.py
- dtypes.py
- initializers.py
- linear.py
- lora.py
- normalization.py
- recurrent.py
- stochastic.py
- requirements.txt
- run-all-examples.bash
- __init__.py
- ema.py
- metrics.py
- optimizer.py
- __init__.py
- autodiff.py
- compilation.py
- general.py
- iteration.py
- transforms.py
- __init__.py
- compat.py
- deprecations.py
- extract.py
- filterlib.py
- graph.py
- graphlib.py
- helpers.py
- ids.py
- module.py
- proxy_caller.py
- pytreelib.py
- README.md
- reprlib.py
- rnglib.py
- spmd.py
- statelib.py
- summary.py
- tracers.py
- traversals.py
- variablelib.py
- visualization.py
- .git-blame-ignore-revs
- __init__.py
- benchmark.py
- __init__.py
- checkpoints.py
- common_utils.py
- dynamic_scale.py
- early_stopping.py
- lr_schedule.py
- orbax_utils.py
- prefetch_iterator.py
- train_state.py
- __init__.py
- configurations.py
- cursor.py
- errors.py
- ids.py
- io.py
- jax_utils.py
- py.typed
- serialization.py
- struct.py
- traceback_util.py
- traverse_util.py
- typing.py
- version.py
- __init__.py
- flaxlib_cpp.pyi
- lib.cc
- .gitignore
- Cargo.lock
- Cargo.toml
- CMakeLists.txt
- LICENSE
- pyproject.toml
- README.md
- uv.lock
- flax_logo.png
- flax_logo.svg
- flax_logo_250px.png
- flax_logo_500px.png
- core_attention_test.py
- core_auto_encoder_test.py
- core_big_resnets_test.py
- core_custom_vjp_test.py
- core_dense_test.py
- core_flow_test.py
- core_resnet_test.py
- core_scan_test.py
- core_tied_autoencoder_test.py
- core_vmap_test.py
- core_weight_std_test.py
- core_frozen_dict_test.py
- core_lift_test.py
- core_meta_test.py
- core_scope_test.py
- initializers_test.py
- kw_only_dataclasses_test.py
- linen_activation_test.py
- linen_attention_test.py
- linen_batch_apply_test.py
- linen_combinators_test.py
- linen_dtypes_test.py
- linen_linear_test.py
- linen_meta_test.py
- linen_module_test.py
- linen_recurrent_test.py
- linen_test.py
- linen_transforms_test.py
- partitioning_test.py
- summary_test.py
- module_test.py
- wrappers_test.py
- attention_test.py
- conv_test.py
- embed_test.py
- linear_test.py
- lora_test.py
- normalization_test.py
- recurrent_test.py
- stochastic_test.py
- __init__.py
- containers_test.py
- ema_test.py
- filters_test.py
- graph_utils_test.py
- helpers_test.py
- ids_test.py
- integration_test.py
- metrics_test.py
- module_test.py
- mutable_array_test.py
- nnx_compat_test.py
- optimizer_test.py
- partitioning_test.py
- rngs_test.py
- spmd_test.py
- state_test.py
- summary_test.py
- test_traversals.py
- transforms_test.py
- variable_test.py
- checkpoints_test.py
- colab_tpu_jax_version.ipynb
- configurations_test.py
- cursor_test.py
- download_dataset_metadata.sh
- early_stopping_test.py
- flaxlib_test.py
- import_test.ipynb
- io_test.py
- jax_utils_test.py
- pickle_test.py
- run_all_tests.sh
- serialization_test.py
- struct_test.py
- tensorboard_test.py
- traceback_util_test.py
- traverse_util_test.py
- .git-blame-ignore-revs
- .gitignore
- .pre-commit-config.yaml
- .readthedocs.yml
- AUTHORS
- CHANGELOG.md
- contributing.md
- LICENSE
- nnx.py
- pylintrc
- pyproject.toml
- README.md
# Installation Guide
git clone https://github.com/google/flax
Downloads the entire project code from GitHub to your computer.
cd flax
Moves into the project folder you just downloaded.
2. Official Install Script
Easy Recommended- Python 3 Python is required to use pip.
pip install flax
Installs the package published on PyPI directly β no need to clone the source.
pip install --upgrade git+https://github.com/google/flax.git
Installs the package published on PyPI directly β no need to clone the source.
pip install "flax[all]"
Installs the package published on PyPI directly β no need to clone the source.
Pulled directly from this repo's README.
3. CMake
Mediumcd flaxlib_src
This project's files live in a subfolder, so move into it first.
mkdir build && cd build
Creates a folder to hold the build output and moves into it.
cmake ..
Analyzes the source code and generates build configuration files (must be run inside the build folder).
make
Compiles the code based on the generated build configuration to produce an executable.
4. Python
Easypip install flax
Installs the package published on PyPI directly β no need to clone the source.
pip install --upgrade git+https://github.com/google/flax.git
Installs the package published on PyPI directly β no need to clone the source.
pip install "flax[all]"
Installs the package published on PyPI directly β no need to clone the source.
Pulled directly from this repo's README.
5. Rust
Medium- Git Needed to download the project code from GitHub.
- Rust (rustup) Installing via rustup also installs cargo.
cd flaxlib_src
This project's files live in a subfolder, so move into it first.
cargo build --release
Compiles the Rust project.
cargo run
Builds and then immediately runs the program.
6. Make
Medium- Git Needed to download the project code from GitHub.
- Make Usually pre-installed on Linux/macOS. On Windows, install separately (e.g. via MSYS2 or WSL).
cd docs_nnx
This project's files live in a subfolder, so move into it first.
make
Compiles the code based on the generated build configuration to produce an executable.
