Update MaxText dependencies#

Introduction#

This document provides a guide to updating dependencies in MaxText using the seed-env tool. seed-env helps generate deterministic and reproducible Python environments by creating fully-pinned requirements.txt files from a base set of requirements.

Please keep dependencies updated throughout development. This will allow each commit to work properly from both a feature and dependency perspective. We will periodically upload commits to PyPI for stable releases. But it is also critical to keep dependencies in sync for users installing MaxText from source.

Overview of the process#

To update dependencies, you will follow these general steps:

  1. Modify base requirements: Update the desired dependencies in src/dependencies/requirements/base_requirements/requirements.txt or the hardware-specific pre-training files (src/dependencies/requirements/base_requirements/tpu-requirements.txt, src/dependencies/requirements/base_requirements/cuda12-requirements.txt) or the post-training files (src/dependencies/requirements/base_requirements/tpu-post-train-requirements.txt).

  2. Find the JAX build commit hash: The dependency generation process is pinned to a specific nightly build of JAX. You need to find the commit hash for the desired JAX build.

  3. Generate the requirement files: Run src/dependencies/scripts/generate_requirements.sh, which internally invokes seed-env to produce fully-pinned requirements files.

  4. Verify the new dependencies: Test the new dependencies to ensure the project installs and runs correctly.

The following sections provide detailed instructions for each step.

Step 0: Install seed-env#

First, you need to install the seed-env command-line tool. We recommend installing uv first following uv’s official installation instructions and then using it to install seed-env:

uv venv --python 3.12 --seed seed_venv
source seed_venv/bin/activate
uv pip install seed-env

Alternatively, follow the instructions in the seed-env repository if you want to build seed-env from source.

Step 1: Modify base requirements#

Update the desired dependencies in src/dependencies/requirements/base_requirements/requirements.txt or the hardware-specific pre-training files (src/dependencies/requirements/base_requirements/tpu-requirements.txt, src/dependencies/requirements/base_requirements/cuda12-requirements.txt) or the post-training files (src/dependencies/requirements/base_requirements/tpu-post-train-requirements.txt).

Put a version floor that only the pre-training environments can satisfy in the hardware-specific pre-training files, not in requirements.txt. The post-training requirements include requirements.txt and are resolved against their own JAX seed before the post-training overrides apply, so a floor that needs a newer JAX than that seed makes the post-training regeneration fail. For example, flax>=0.12.10 needs jax>=0.11.1, but the post-training seed pins jaxlib==0.11.0.

Step 2: Find the JAX build commit hash#

The dependency generation process is pinned to a specific nightly build of JAX. You need to find the commit hash for the desired JAX build from JAX build/ folder and copy its full commit hash.

Step 3: Generate the requirements files#

Next, run generate_requirements.sh to generate the new requirements files. This script wraps the seed-env CLI and handles exporting the lock, and applying any overrides. You will need to do this separately for the TPU and GPU environments.

TPU Pre-Training#

Note: The current src/dependencies/requirements/generated_requirements/tpu-requirements.txt in the repository was generated using JAX build commit hash: ab4c9b943c70bcb42baf4d379036a19c8aa2689d. When regenerating the requirements, either use the same commit hash or update this hash if you use a different one.

If you have made changes to TPU pre-training dependencies in src/dependencies/requirements/base_requirements/tpu-requirements.txt, you need to regenerate the pinned pre-training requirements in generated_requirements/ directory. Run the following command, replacing <jax-build-commit-hash> with the hash you copied in the previous step:

bash src/dependencies/scripts/generate_requirements.sh \
--base-requirements src/dependencies/requirements/base_requirements/tpu-requirements.txt \
--generated-requirements tpu-requirements.txt \
--override-requirements src/dependencies/extra_deps/tpu_overrides.txt \
--seed-commit ab4c9b943c70bcb42baf4d379036a19c8aa2689d

# Copy generated requirements to src/dependencies/requirements/generated_requirements
mv generated_artifacts/python3_12/tpu-requirements.txt \
  src/dependencies/requirements/generated_requirements/tpu-requirements.txt

TPU Post-Training#

Note: The current src/dependencies/requirements/generated_requirements/tpu-post-train-requirements.txt in the repository was generated using JAX build commit hash: 52d5cb3893727451d0f695bfa6f8206cdd9492c7. When regenerating the requirements, either use the same commit hash or update this hash if you use a different one.

If you have made changes to the post-training dependencies in src/dependencies/requirements/base_requirements/tpu-post-train-requirements.txt, you need to regenerate the pinned post-training requirements in generated_requirements/ directory. Run the following command, replacing <jax-build-commit-hash> with the hash you copied in the previous step:

bash src/dependencies/scripts/generate_requirements.sh \
--base-requirements src/dependencies/requirements/base_requirements/tpu-post-train-requirements.txt \
--generated-requirements tpu-post-train-requirements.txt \
--override-requirements src/dependencies/extra_deps/tpu_post_train_overrides.txt \
--seed-commit 52d5cb3893727451d0f695bfa6f8206cdd9492c7

# Copy generated requirements to src/dependencies/requirements/generated_requirements
mv generated_artifacts/python3_12/tpu-post-train-requirements.txt \
  src/dependencies/requirements/generated_requirements/tpu-post-train-requirements.txt

GPU Pre-Training#

Note: The current src/dependencies/requirements/generated_requirements/cuda12-requirements.txt in the repository was generated using JAX build commit hash: ab4c9b943c70bcb42baf4d379036a19c8aa2689d. When regenerating the requirements, either use the same commit hash or update this hash if you use a different one.

If you have made changes to the GPU pre-training dependencies in src/dependencies/requirements/base_requirements/cuda12-requirements.txt, you need to regenerate the pinned pre-training requirements in generated_requirements/ directory. Run the following command, replacing <jax-build-commit-hash> with the hash you copied in the previous step:

bash src/dependencies/scripts/generate_requirements.sh \
--base-requirements src/dependencies/requirements/base_requirements/cuda12-requirements.txt \
--generated-requirements cuda12-requirements.txt \
--seed-commit ab4c9b943c70bcb42baf4d379036a19c8aa2689d \
--override-requirements src/dependencies/extra_deps/cuda12_overrides.txt \
--hardware cuda12

# Copy generated requirements to src/dependencies/requirements/generated_requirements
mv generated_artifacts/python3_12/cuda12-requirements.txt \
  src/dependencies/requirements/generated_requirements/cuda12-requirements.txt

Decoupled mode#

generated_requirements/decoupled-requirements.txt is derived from the GPU pre-training requirements above: the same packages at the same versions, minus the Google Cloud clients and the accelerator wheels. Decoupled mode is only exercised when those packages are genuinely absent, so this file is what CI installs for the decoupled test run. It carries no seed-env step of its own, it is a filter over a lock that seed-env already produced, which is why regenerating it needs no network. Pass --source to derive the same thing from another hardware lock. Regenerate it after updating any requirements file:

python3 src/dependencies/scripts/generate_decoupled_requirements.py

The decoupled-requirements pre-commit hook fails when a change under src/dependencies/requirements/ leaves this file stale.

Optional TensorFlow and JetStream dependencies#

Optional TensorFlow, SeqIO, and JetStream dependencies are listed with pinned versions in src/dependencies/extra_deps/tf_requirements.txt and installed with --no-deps when --with-tf is passed to install_pre_train_extra_deps.py.

To update src/dependencies/extra_deps/tf_requirements.txt when upgrading TensorFlow or JetStream, run src/dependencies/scripts/generate_tf_requirements.py. This script invokes uv pip compile using tpu-requirements.txt as a strict constraint file (-c) to prevent conflicts with core MaxText dependencies and writes only the pinned delta to src/dependencies/extra_deps/tf_requirements.txt:

python3 src/dependencies/scripts/generate_tf_requirements.py \
  --tensorflow-version 2.20.0 \
  --tensorflow-text-version 2.20.1 \
  --jetstream-commit 29329e8e73820993f77cfc8efe34eb2a73f5de98

Step 4: Verify the new dependencies#

Finally, test that the new dependencies install correctly and that MaxText runs as expected.

  1. Install MaxText and dependencies: For instructions on installing MaxText on your VM, please refer to the official documentation.

  2. Run tests: Run MaxText tests to ensure there are no regressions.