Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/dependencies/extra_deps/tpu_overrides.txt
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
datasets>=4.8.5
fsspec==2026.2.0
gcsfs==2026.2.0
orbax-checkpoint>=0.12.1
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,12 @@
absl-py>=2.4.0
aiofiles>=25.1.0
aiohappyeyeballs>=2.6.2
aiohttp>=3.13.5
aiohttp>=3.14.1
aiosignal>=1.4.0
annotated-doc>=0.0.4
annotated-types>=0.7.0
antlr4-python3-runtime>=4.9.3
anyio>=4.13.0
anyio>=4.14.1
aqtp>=0.9.0
array-record>=0.8.3
astroid>=4.0.4
Expand All @@ -22,29 +22,29 @@ certifi>=2026.2.25
cffi>=2.0.0 ; platform_python_implementation != 'PyPy'
cfgv>=3.5.0
charset-normalizer>=3.4.7
chex>=0.1.91
click>=8.4.0
chex>=0.1.92
click>=8.4.2
cloud-accelerator-diagnostics>=0.1.1
cloudpickle>=3.1.2
clu>=0.0.12
colorama>=0.4.6
contourpy>=1.3.3
cryptography>=48.0.0
cryptography>=49.0.0
cycler>=0.12.1
datasets>=4.8.5
datasets>=5.0.0
decorator>=5.3.1
dill>=0.4.1
distlib>=0.4.0
distlib>=0.4.3
distro>=1.9.0
dm-tree>=0.1.10
docstring-parser>=0.18.0
drjax>=0.1.4
drjax>=0.2.0
editdistance>=0.8.1
einops>=0.8.2
einshape>=1.0
etils>=1.14.0
execnet>=2.1.2
fastapi>=0.136.1
fastapi>=0.138.0
filelock>=3.28.0
flatbuffers>=25.12.19
flax>=0.12.7
Expand All @@ -53,39 +53,39 @@ frozenlist>=1.8.0
fsspec>=2026.2.0
gast>=0.7.0
gcsfs>=2026.2.0
google-api-core>=2.30.3
google-api-python-client>=2.196.0
google-auth>=2.53.0
google-api-core>=2.31.0
google-api-python-client>=2.197.0
google-auth>=2.55.0
google-auth-httplib2>=0.4.0
google-auth-oauthlib>=1.4.0
google-cloud-aiplatform>=1.153.1
google-cloud-appengine-logging>=1.9.0
google-cloud-audit-log>=0.5.0
google-cloud-bigquery>=3.41.0
google-cloud-aiplatform>=1.158.0
google-cloud-appengine-logging>=1.10.0
google-cloud-audit-log>=0.6.0
google-cloud-bigquery>=3.42.1
google-cloud-core>=2.6.0
google-cloud-logging>=3.15.0
google-cloud-mldiagnostics>=1.0.2
google-cloud-monitoring>=2.30.0
google-cloud-resource-manager>=1.17.0
google-cloud-storage>=3.10.1
google-cloud-storage-control>=1.11.0
google-cloud-logging>=3.16.0
google-cloud-mldiagnostics>=1.0.3
google-cloud-monitoring>=2.31.0
google-cloud-resource-manager>=1.18.0
google-cloud-storage>=3.12.0
google-cloud-storage-control>=1.12.0
google-crc32c>=1.8.0
google-genai>=2.4.0
google-genai>=2.10.0
google-pasta>=0.2.0
google-resumable-media>=2.9.0
google-resumable-media>=2.10.0
googleapis-common-protos>=1.75.0
grain>=0.2.16
grain>=0.2.18
grpc-google-iam-v1>=0.14.4
grpcio>=1.80.0
grpcio-status>=1.80.0
gviz-api>=1.10.0
h11>=0.16.0
h5py>=3.14.0
hf-xet>=1.5.0 ; platform_machine == 'AMD64' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64'
hf-xet>=1.5.1 ; platform_machine == 'AMD64' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64'
httpcore>=1.0.9
httplib2>=0.31.2
httpx>=0.28.1
huggingface-hub>=1.15.0
huggingface-hub>=1.20.1
humanize>=4.15.0
hypothesis>=6.142.1
identify>=2.6.19
Expand All @@ -96,7 +96,7 @@ iniconfig>=2.3.0
isort>=8.0.1
jax>=0.10.0
jaxlib>=0.10.0
jaxtyping>=0.3.9
jaxtyping>=0.3.11
jinja2>=3.1.6
jsonlines>=4.0.0
keras>=3.14.0
Expand All @@ -114,9 +114,9 @@ mccabe>=0.7.0
mdurl>=0.1.2
ml-collections>=1.1.0
ml-dtypes>=0.5.4
ml-goodput-measurement>=0.0.16
ml-goodput-measurement>=0.2.0
mpmath>=1.3.0
msgpack>=1.1.2
msgpack>=1.2.1
msgspec>=0.21.1
multidict>=6.7.1
multiprocess>=0.70.19
Expand All @@ -129,24 +129,26 @@ nodeenv>=1.10.0
numpy>=2.0.2
numpy-typing-compat>=20251206.2.0
nvidia-cuda-cccl>=13.2.75
nvidia-ml-py>=13.610.43
oauthlib>=3.3.1
omegaconf>=2.3.0
opentelemetry-api>=1.42.0
omegaconf>=2.3.1
opentelemetry-api>=1.43.0
opt-einsum>=3.4.0
optax>=0.2.8
optree>=0.19.0
optype>=0.17.0
orbax-checkpoint>=0.11.39
orbax-checkpoint>=0.12.1
packaging>=26.1
pandas>=3.0.3
parameterized>=0.9.0
pathspec>=1.1.1
pathwaysutils>=0.1.8
pathwaysutils>=0.1.9
pillow>=12.2.0
platformdirs>=4.9.6
platformdirs>=4.10.0
pluggy>=1.6.0
portpicker>=1.6.0
pre-commit>=4.6.0
prometheus-client>=0.20.0
promise>=2.3
propcache>=0.5.2
proto-plus>=1.28.0
Expand All @@ -164,22 +166,24 @@ pyelftools>=0.32
pyglove>=0.4.5
pygments>=2.20.0
pyink>=25.12.0
pylint>=4.0.5
pylint>=4.0.6
pynvml>=13.0.1
pyopenssl>=26.3.0
pyparsing>=3.3.2
pyproject-hooks>=1.2.0
pytest>=8.4.2
pytest-xdist>=3.8.0
python-dateutil>=2.9.0.post0
python-discovery>=1.3.1
python-discovery>=1.4.2
pytokens>=0.4.1
pytype>=2024.10.11
pyyaml>=6.0.3
qwix>=0.1.6
qwix>=0.1.8
regex>=2026.5.9
requests>=2.33.1
requests-oauthlib>=2.0.0
rich>=15.0.0
safetensors>=0.7.0
safetensors>=0.8.0
scipy>=1.17.1
scipy-stubs>=1.17.1.4
sentencepiece>=0.2.1
Expand All @@ -191,7 +195,7 @@ simplejson>=4.1.1
six>=1.17.0
sniffio>=1.3.1
sortedcontainers>=2.4.0
starlette>=1.0.0
starlette>=1.3.1
sympy>=1.14.0
tabulate>=0.10.0
tenacity>=9.1.4
Expand All @@ -201,18 +205,18 @@ tensorboard-plugin-profile>=2.13.0
tensorboardx>=2.6.5
tensorflow>=2.20.0
tensorflow-datasets>=4.9.10
tensorflow-metadata>=1.17.3
tensorflow-metadata>=1.21.0
tensorflow-text>=2.20.1
tensorstore>=0.1.82
tensorstore>=0.1.84
termcolor>=3.3.0
tiktoken>=0.13.0
tokamax>=0.0.12
tokenizers>=0.22.2
toml>=0.10.2
tomlkit>=0.15.0
toolz>=1.1.0
tqdm>=4.66.3
transformers>=5.9.0
tqdm>=4.68.3
transformers>=5.12.1
treescope>=0.1.10
typeguard>=2.13.3
typer>=0.25.1
Expand All @@ -221,15 +225,15 @@ typing-inspection>=0.4.2
tzdata>=2026.2 ; sys_platform == 'emscripten' or sys_platform == 'win32'
uritemplate>=4.2.0
urllib3>=2.6.3
uvicorn>=0.47.0
uvicorn>=0.49.0
uvloop>=0.22.1
virtualenv>=21.3.3
virtualenv>=21.5.1
wadler-lindig>=0.1.7
websockets>=16.0
werkzeug>=3.1.8
wheel>=0.46.3
wrapt>=2.1.2
xxhash>=3.7.0
wrapt>=2.2.2
xxhash>=3.7.1
yarl>=1.24.2
zipp>=3.23.1
zstandard>=0.25.0
33 changes: 7 additions & 26 deletions src/maxtext/common/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -443,17 +443,6 @@ def create_orbax_checkpoint_manager(
logger=orbax_logger,
)

# Use Colocated Python checkpointing optimization (Single Controller only).
if enable_single_controller and colocated_python_checkpointing:
max_logging.log("Registering colocated python array handler")
checkpointing_impl = ocp.pathways.CheckpointingImpl.from_options(
use_colocated_python=True,
)
ocp.pathways.register_type_handlers(
use_single_replica_array_handler=enable_single_replica_ckpt_restoring,
checkpointing_impl=checkpointing_impl,
)

max_logging.log("Checkpoint manager created!")
return manager

Expand Down Expand Up @@ -500,6 +489,7 @@ def create_orbax_emergency_replicator_checkpoint_manager(
local_checkpoint_dir: str,
save_interval_steps: int,
global_mesh: jax.sharding.Mesh,
colocated_python_checkpointing: bool = False,
):
"""Returns an emergency replicator checkpoint manager."""
flags.FLAGS.experimental_orbax_use_distributed_process_id = True
Expand All @@ -509,6 +499,7 @@ def create_orbax_emergency_replicator_checkpoint_manager(
epath.Path(local_checkpoint_dir),
options=emergency_replicator_checkpoint_manager.ReplicatorCheckpointManagerOptions(
save_interval_steps=save_interval_steps,
use_colocated_python=colocated_python_checkpointing,
),
global_mesh=global_mesh,
)
Expand Down Expand Up @@ -838,9 +829,7 @@ def map_to_pspec(data):
(EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager),
):
checkpoint_path = str(checkpoint_manager.directory / str(step) / "items")
with handle_checkpoint_mismatch(
"restore NNX checkpoint", checkpoint_path
):
with handle_checkpoint_mismatch("restore NNX checkpoint", checkpoint_path):
Comment thread
SujeethJinesh marked this conversation as resolved.
restored_nnx = _load_linen_checkpoint_into_nnx(
checkpoint_path,
abstract_unboxed_pre_state,
Expand Down Expand Up @@ -876,9 +865,7 @@ def map_to_pspec(data):
EmergencyReplicatorCheckpointManager,
),
):
restored = checkpoint_manager.restore(
step, args=Composite(state=checkpoint_args)
).state
restored = checkpoint_manager.restore(step, args=Composite(state=checkpoint_args)).state
_assert_no_shaped_dtype_struct(restored)
return (
restored,
Expand Down Expand Up @@ -906,9 +893,7 @@ def map_to_pspec(data):
# Case 3: Default/Fallback case.
# This case acts as a wildcard ('_') and matches if none of the preceding cases were met.
case _:
restored = checkpoint_manager.restore(
step, args=Composite(items=checkpoint_args)
)
restored = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args))
_assert_no_shaped_dtype_struct(restored)
return (restored, None)

Expand All @@ -918,9 +903,7 @@ def map_to_pspec(data):
else:
params = abstract_unboxed_pre_state.params

with handle_checkpoint_mismatch(
"load parameters", load_parameters_from_path
):
with handle_checkpoint_mismatch("load parameters", load_parameters_from_path):
restored_params = load_params_from_path(
load_parameters_from_path,
params,
Expand All @@ -932,9 +915,7 @@ def map_to_pspec(data):
return None, restored_params
elif load_full_state_from_path != "":
max_logging.log(f"Loading full state from path: {load_full_state_from_path}")
with handle_checkpoint_mismatch(
"load full state", load_full_state_from_path
):
with handle_checkpoint_mismatch("load full state", load_full_state_from_path):
restored_state = _load_full_state_from_path(
path=load_full_state_from_path,
abstract_unboxed_pre_state=abstract_unboxed_pre_state,
Expand Down
12 changes: 12 additions & 0 deletions src/maxtext/utils/max_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,6 +246,17 @@ def maybe_initialize_jax_distributed_system(raw_keys):
return
if raw_keys["enable_single_controller"]:
max_logging.log("Skipping jax distributed system since its not needed for single controller.")
if raw_keys["enable_multi_tier_checkpointing"]:
max_logging.log("Initializing multi-tier checkpointing for single controller...")
initialize_multi_tier_checkpointing(
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"],
run_name=raw_keys["run_name"],
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
data_parallelism=raw_keys["mtc_data_parallelism"],
num_slices=raw_keys["num_slices"],
use_colocated_python=True,
)
return
if jax.distributed.is_initialized():
max_logging.log("Jax distributed system is already initialized.")
Expand Down Expand Up @@ -290,6 +301,7 @@ def maybe_initialize_jax_distributed_system(raw_keys):
run_name=raw_keys["run_name"],
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
data_parallelism=raw_keys["mtc_data_parallelism"],
num_slices=raw_keys["num_slices"],
)
max_logging.log("Jax distributed system initialized on TPUs for multi-tier checkpointing!")
elif raw_keys["enable_checkpointing"] and raw_keys["compile_topology_num_slices"] == -1:
Expand Down
13 changes: 13 additions & 0 deletions src/maxtext/utils/train_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
import subprocess
import jax
import functools
import orbax.checkpoint.pathways as ocp_pathways
from functools import partial

from flax import nnx
Expand Down Expand Up @@ -55,6 +56,7 @@ def create_checkpoint_manager(config, mesh, init_state_fn):
config.local_checkpoint_directory,
config.local_checkpoint_period,
mesh,
config.colocated_python_checkpointing,
)
elif config.enable_emergency_checkpoint:
abstract_state, _, _ = maxtext_utils.get_abstract_state(config, mesh, init_state_fn, is_training=True)
Expand Down Expand Up @@ -97,6 +99,17 @@ def create_checkpoint_manager(config, mesh, init_state_fn):
config.checkpoint_todelete_full_path,
)

# Use Colocated Python checkpointing dispatchers optimization (Single Controller only).
if checkpoint_manager is not None and config.enable_single_controller and config.colocated_python_checkpointing:
max_logging.log("Registering colocated python array handler")
checkpointing_impl = ocp_pathways.CheckpointingImpl.from_options(
use_colocated_python=True,
)
ocp_pathways.register_type_handlers(
use_single_replica_array_handler=config.enable_single_replica_ckpt_restoring,
checkpointing_impl=checkpointing_impl,
)

return checkpoint_manager


Expand Down
Loading
Loading