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
// repository documentation
Was this content helpful?
(0 ratings)
