AXLearn: Apple's JAX-Based Framework for Large-Scale Model Training
An Extensible Deep Learning Library
At a glance
- What is it?
- AXLearn is an open-source Python library from Apple built on JAX and XLA, designed for training large deep learning models through a composable configuration system that scales to hundreds of billions of parameters across thousands of accelerators.
- Who is it for?
- AXLearn is built for teams that train large models at scale on cloud infrastructure and want a configuration system that separates model structure from weight allocation. Engineers who build on PyTorch or need a framework with stable releases and a guarantee of API continuity should look elsewhere: the README states the API is subject to change.
- Can I use it commercially?
- Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
- Is it still maintained?
- Yes. The repository last received commits 84 days ago.
- What is it written in?
- Mainly Python, according to GitHub's language statistics.
Answers come from the project's GitHub data, last synced on September 30, 2026, and from our analysis. They are not legal advice.
Editorial analysis
What Problem AXLearn Addresses and Who It Is For
Training a model with hundreds of billions of parameters is not just a compute problem. It is a software engineering problem: how do you describe a model composed of dozens of reusable building blocks, integrate it with data pipelines, and reproduce its configuration reliably across many accelerators? AXLearn addresses these questions with a library built on top of JAX and XLA.
The primary audience is research engineers who work on large models in cloud environments. The README describes support for training models with up to hundreds of billions of parameters across thousands of accelerators at high utilisation. The library includes tools for deploying and managing jobs and data on public clouds. It also supports NLP, computer vision, and speech recognition workloads, and the README states it contains baseline configurations for training state-of-the-art models.
AXLearn is not aimed at practitioners who want a quick way to fine-tune an existing checkpoint with minimal code. The configuration system and the dependency stack both assume a team with the engineering resources to set up and maintain a large training pipeline.
Object-Oriented Configuration and Modular Composition
The design centre of AXLearn is its configuration system. The README describes an object-oriented approach to the software engineering challenges that arise from building, iterating, and maintaining models. Rather than defining a model imperatively, operators compose it from reusable building blocks through a configuration object. That object describes the component structure before any weights are allocated.
This separation has practical consequences. A configuration for a transformer encoder can be built by composing attention, feedforward, and normalisation building blocks, each of which can be swapped or parameterised independently. The system also allows integration with Flax and Hugging Face transformers components, so teams that have existing Flax modules or Hugging Face model weights can incorporate them into an AXLearn training pipeline.
A comparison with Flax, which the README names as an integration target, illustrates the design choice. Flax uses a module-centric approach where each component is a Python class that manages its own parameters. AXLearn's configuration system keeps the structural description separate from instantiation. For large composite models, that separation makes it easier to reason about the full configuration as a first-class artifact that can be versioned, diffed, or generated programmatically.
Scaling with GSPMD and Global Computation
AXLearn builds on GSPMD, a partitioning system documented in an arXiv paper that the README cites. The README describes the outcome as a global computation paradigm: operators write code that describes computation on a virtual global computer rather than on a per-accelerator basis. The underlying XLA compiler handles the sharding and communication between accelerators.
This matters at large scales because manually managing tensor sharding across thousands of devices is error-prone. The global view abstracts that complexity. A training run that targets 8 accelerators and one that targets 8,000 can use the same model code; the sharding configuration changes, not the model definition itself.
The library targets public cloud infrastructure and ships with tooling for deploying and managing training jobs and data. The repository includes a Dockerfile at the top level for containerised environments. The Dockerfile uses Ubuntu 24.04 as a base, installs Python 3.12, and sets up a virtual environment. The exact dependency resolution for a production training job would go through the pyproject.toml, which pins JAX and XLA at 0.8.3, TensorFlow at 2.19.1.3, and a number of other packages at exact versions to ensure reproducibility across workers.
Installation and the Dependency Stack
The README directs users to RELEASE.md for installation instructions and the PyPI release process. That file is not included in the repository materials provided here, so no pip install command is reproduced. The pyproject.toml names the package `axlearn` and requires Python exactly 3.12. The core dependency group includes JAX 0.8.3, JAXlib 0.8.3, TensorFlow 2.19.1.3, optax 0.2.6, seqio, and several other packages pinned at exact versions.
The version pinning is intentional and documented in the pyproject.toml with inline comments explaining why certain packages are pinned: for example, pyarrow is capped below 21.0.0 to avoid a breaking type change, and tensorflow-io is pinned at 0.37.3 to avoid a pure-virtual-method crash in earlier patch releases. This means installing AXLearn into an existing Python environment with different versions of these packages is likely to produce dependency conflicts.
The practical path for most users will be a dedicated virtual environment or a container based on the provided Dockerfile. The repository also includes a Bazel build file (.bazelrc, .bazelversion, BUILD.bazel, MODULE.bazel) for teams that use Bazel as their build system.
Supported Applications and Baseline Configurations
The README identifies three application areas: natural language processing, computer vision, and speech recognition. For each, the library includes baseline configurations for training state-of-the-art models. The README does not enumerate specific model architectures or name specific baseline configurations. The docs/ directory exists in the repository, with guides covering getting started, core concepts, CLI usage, and infrastructure.
The repository structure includes an axlearn/ directory for the library itself. The configuration system's design principle is that any model component should be reusable and composable. The README describes this as the central software engineering goal: building, iterating, and maintaining models becomes more manageable when components are expressed as configuration objects rather than hardwired imperative code.
The Concepts guide at docs/02-concepts.md is cited in the README as the place to understand the core components and design principles. Anyone planning to extend or customise AXLearn should read it before writing new components, since the configuration system's rules govern how components discover and depend on each other.
Limitations and Situations Where AXLearn Is Not the Right Tool
AXLearn is not the right tool for engineers who work primarily in PyTorch. The library is built entirely on JAX and XLA. PyTorch checkpoints and PyTorch-native operators do not work in an AXLearn training loop without conversion.
The dependency pinning that ensures reproducibility across large training runs also makes AXLearn difficult to install alongside other libraries that have their own version requirements. A project that already depends on a different version of JAX or TensorFlow will face conflicts. The README explicitly states that the library is under active development and the API is subject to change, which means code written against the current API may break in a future release. There are no GitHub releases in the repository, so version selection must be done by commit hash.
Teams that want to fine-tune a pre-trained model in a few lines of code and do not need to train at scale have simpler alternatives. Hugging Face Accelerate, which is a widely adopted training framework for PyTorch and JAX, is one such option; it focuses on multi-device training without the large configuration-system overhead that AXLearn imposes.
Maintenance Status and Licensing
The last push to the main branch was on 2026-07-08. The repository has no GitHub releases. The README carries a notice at the top that the library is under active development and the API is subject to change. The CHANGELOG.md, CONTRIBUTING.md, and CODEOWNERS files are present at the top level.
The library is Apache-2.0 licensed. That licence permits commercial use, modification, and redistribution, with attribution requirements. The dependency stack includes TensorFlow, which carries its own licence terms; teams that build a product on AXLearn should review the full dependency list in pyproject.toml to confirm that the licences of all pinned packages are compatible with their deployment context.
The repository was created by Apple and the CODEOWNERS file indicates Apple employees maintain the primary review responsibilities. External contributions are accepted according to CONTRIBUTING.md.
Editorial conclusion
AXLearn is built for teams that train large models at scale on cloud infrastructure and want a configuration system that separates model structure from weight allocation. Engineers who build on PyTorch or need a framework with stable releases and a guarantee of API continuity should look elsewhere: the README states the API is subject to change. Before adopting it, review the dependency list in pyproject.toml, which pins JAX 0.8.3, TensorFlow 2.19.1.3, and several other packages at exact versions, and verify compatibility with your accelerator environment.
Frequently asked questions
What Python version does AXLearn require?
The pyproject.toml specifies Python exactly 3.12. Other Python versions, including 3.11, are not listed as supported.
Does AXLearn work with PyTorch models?
AXLearn is built on JAX and XLA, not PyTorch. The README names Flax and Hugging Face transformers as integration targets, but does not mention PyTorch compatibility.
Where are the installation instructions for AXLearn?
The README refers to RELEASE.md in the repository for the PyPI release process, version management, and installation instructions.
Official sources
Add this badge to your README
If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.
[](https://hysenlabs.com/projects/apple-axlearn)