Skip to content

Commit 7549b91

Browse files
Merge pull request AI-Hypercomputer#3952 from AI-Hypercomputer:sujinesh/mtc_pathways_support
PiperOrigin-RevId: 938794444
2 parents df55b86 + 73cd6ad commit 7549b91

9 files changed

Lines changed: 188 additions & 95 deletions

File tree

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
11
datasets>=4.8.5
22
fsspec==2026.2.0
33
gcsfs==2026.2.0
4+
orbax-checkpoint>=0.12.1

src/dependencies/requirements/generated_requirements/tpu-requirements.txt

Lines changed: 52 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,12 @@
44
absl-py>=2.4.0
55
aiofiles>=25.1.0
66
aiohappyeyeballs>=2.6.2
7-
aiohttp>=3.13.5
7+
aiohttp>=3.14.1
88
aiosignal>=1.4.0
99
annotated-doc>=0.0.4
1010
annotated-types>=0.7.0
1111
antlr4-python3-runtime>=4.9.3
12-
anyio>=4.13.0
12+
anyio>=4.14.1
1313
aqtp>=0.9.0
1414
array-record>=0.8.3
1515
astroid>=4.0.4
@@ -22,29 +22,29 @@ certifi>=2026.2.25
2222
cffi>=2.0.0 ; platform_python_implementation != 'PyPy'
2323
cfgv>=3.5.0
2424
charset-normalizer>=3.4.7
25-
chex>=0.1.91
26-
click>=8.4.0
25+
chex>=0.1.92
26+
click>=8.4.2
2727
cloud-accelerator-diagnostics>=0.1.1
2828
cloudpickle>=3.1.2
2929
clu>=0.0.12
3030
colorama>=0.4.6
3131
contourpy>=1.3.3
32-
cryptography>=48.0.0
32+
cryptography>=49.0.0
3333
cycler>=0.12.1
34-
datasets>=4.8.5
34+
datasets>=5.0.0
3535
decorator>=5.3.1
3636
dill>=0.4.1
37-
distlib>=0.4.0
37+
distlib>=0.4.3
3838
distro>=1.9.0
3939
dm-tree>=0.1.10
4040
docstring-parser>=0.18.0
41-
drjax>=0.1.4
41+
drjax>=0.2.0
4242
editdistance>=0.8.1
4343
einops>=0.8.2
4444
einshape>=1.0
4545
etils>=1.14.0
4646
execnet>=2.1.2
47-
fastapi>=0.136.1
47+
fastapi>=0.138.0
4848
filelock>=3.28.0
4949
flatbuffers>=25.12.19
5050
flax>=0.12.7
@@ -53,39 +53,39 @@ frozenlist>=1.8.0
5353
fsspec>=2026.2.0
5454
gast>=0.7.0
5555
gcsfs>=2026.2.0
56-
google-api-core>=2.30.3
57-
google-api-python-client>=2.196.0
58-
google-auth>=2.53.0
56+
google-api-core>=2.31.0
57+
google-api-python-client>=2.197.0
58+
google-auth>=2.55.0
5959
google-auth-httplib2>=0.4.0
6060
google-auth-oauthlib>=1.4.0
61-
google-cloud-aiplatform>=1.153.1
62-
google-cloud-appengine-logging>=1.9.0
63-
google-cloud-audit-log>=0.5.0
64-
google-cloud-bigquery>=3.41.0
61+
google-cloud-aiplatform>=1.158.0
62+
google-cloud-appengine-logging>=1.10.0
63+
google-cloud-audit-log>=0.6.0
64+
google-cloud-bigquery>=3.42.1
6565
google-cloud-core>=2.6.0
66-
google-cloud-logging>=3.15.0
67-
google-cloud-mldiagnostics>=1.0.2
68-
google-cloud-monitoring>=2.30.0
69-
google-cloud-resource-manager>=1.17.0
70-
google-cloud-storage>=3.10.1
71-
google-cloud-storage-control>=1.11.0
66+
google-cloud-logging>=3.16.0
67+
google-cloud-mldiagnostics>=1.0.3
68+
google-cloud-monitoring>=2.31.0
69+
google-cloud-resource-manager>=1.18.0
70+
google-cloud-storage>=3.12.0
71+
google-cloud-storage-control>=1.12.0
7272
google-crc32c>=1.8.0
73-
google-genai>=2.4.0
73+
google-genai>=2.10.0
7474
google-pasta>=0.2.0
75-
google-resumable-media>=2.9.0
75+
google-resumable-media>=2.10.0
7676
googleapis-common-protos>=1.75.0
77-
grain>=0.2.16
77+
grain>=0.2.18
7878
grpc-google-iam-v1>=0.14.4
7979
grpcio>=1.80.0
8080
grpcio-status>=1.80.0
8181
gviz-api>=1.10.0
8282
h11>=0.16.0
8383
h5py>=3.14.0
84-
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'
84+
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'
8585
httpcore>=1.0.9
8686
httplib2>=0.31.2
8787
httpx>=0.28.1
88-
huggingface-hub>=1.15.0
88+
huggingface-hub>=1.20.1
8989
humanize>=4.15.0
9090
hypothesis>=6.142.1
9191
identify>=2.6.19
@@ -96,7 +96,7 @@ iniconfig>=2.3.0
9696
isort>=8.0.1
9797
jax>=0.10.0
9898
jaxlib>=0.10.0
99-
jaxtyping>=0.3.9
99+
jaxtyping>=0.3.11
100100
jinja2>=3.1.6
101101
jsonlines>=4.0.0
102102
keras>=3.14.0
@@ -114,9 +114,9 @@ mccabe>=0.7.0
114114
mdurl>=0.1.2
115115
ml-collections>=1.1.0
116116
ml-dtypes>=0.5.4
117-
ml-goodput-measurement>=0.0.16
117+
ml-goodput-measurement>=0.2.0
118118
mpmath>=1.3.0
119-
msgpack>=1.1.2
119+
msgpack>=1.2.1
120120
msgspec>=0.21.1
121121
multidict>=6.7.1
122122
multiprocess>=0.70.19
@@ -129,24 +129,26 @@ nodeenv>=1.10.0
129129
numpy>=2.0.2
130130
numpy-typing-compat>=20251206.2.0
131131
nvidia-cuda-cccl>=13.2.75
132+
nvidia-ml-py>=13.610.43
132133
oauthlib>=3.3.1
133-
omegaconf>=2.3.0
134-
opentelemetry-api>=1.42.0
134+
omegaconf>=2.3.1
135+
opentelemetry-api>=1.43.0
135136
opt-einsum>=3.4.0
136137
optax>=0.2.8
137138
optree>=0.19.0
138139
optype>=0.17.0
139-
orbax-checkpoint>=0.11.39
140+
orbax-checkpoint>=0.12.1
140141
packaging>=26.1
141142
pandas>=3.0.3
142143
parameterized>=0.9.0
143144
pathspec>=1.1.1
144-
pathwaysutils>=0.1.8
145+
pathwaysutils>=0.1.9
145146
pillow>=12.2.0
146-
platformdirs>=4.9.6
147+
platformdirs>=4.10.0
147148
pluggy>=1.6.0
148149
portpicker>=1.6.0
149150
pre-commit>=4.6.0
151+
prometheus-client>=0.20.0
150152
promise>=2.3
151153
propcache>=0.5.2
152154
proto-plus>=1.28.0
@@ -164,22 +166,24 @@ pyelftools>=0.32
164166
pyglove>=0.4.5
165167
pygments>=2.20.0
166168
pyink>=25.12.0
167-
pylint>=4.0.5
169+
pylint>=4.0.6
170+
pynvml>=13.0.1
171+
pyopenssl>=26.3.0
168172
pyparsing>=3.3.2
169173
pyproject-hooks>=1.2.0
170174
pytest>=8.4.2
171175
pytest-xdist>=3.8.0
172176
python-dateutil>=2.9.0.post0
173-
python-discovery>=1.3.1
177+
python-discovery>=1.4.2
174178
pytokens>=0.4.1
175179
pytype>=2024.10.11
176180
pyyaml>=6.0.3
177-
qwix>=0.1.6
181+
qwix>=0.1.8
178182
regex>=2026.5.9
179183
requests>=2.33.1
180184
requests-oauthlib>=2.0.0
181185
rich>=15.0.0
182-
safetensors>=0.7.0
186+
safetensors>=0.8.0
183187
scipy>=1.17.1
184188
scipy-stubs>=1.17.1.4
185189
sentencepiece>=0.2.1
@@ -191,7 +195,7 @@ simplejson>=4.1.1
191195
six>=1.17.0
192196
sniffio>=1.3.1
193197
sortedcontainers>=2.4.0
194-
starlette>=1.0.0
198+
starlette>=1.3.1
195199
sympy>=1.14.0
196200
tabulate>=0.10.0
197201
tenacity>=9.1.4
@@ -201,18 +205,18 @@ tensorboard-plugin-profile>=2.13.0
201205
tensorboardx>=2.6.5
202206
tensorflow>=2.20.0
203207
tensorflow-datasets>=4.9.10
204-
tensorflow-metadata>=1.17.3
208+
tensorflow-metadata>=1.21.0
205209
tensorflow-text>=2.20.1
206-
tensorstore>=0.1.82
210+
tensorstore>=0.1.84
207211
termcolor>=3.3.0
208212
tiktoken>=0.13.0
209213
tokamax>=0.0.12
210214
tokenizers>=0.22.2
211215
toml>=0.10.2
212216
tomlkit>=0.15.0
213217
toolz>=1.1.0
214-
tqdm>=4.66.3
215-
transformers>=5.9.0
218+
tqdm>=4.68.3
219+
transformers>=5.12.1
216220
treescope>=0.1.10
217221
typeguard>=2.13.3
218222
typer>=0.25.1
@@ -221,15 +225,15 @@ typing-inspection>=0.4.2
221225
tzdata>=2026.2 ; sys_platform == 'emscripten' or sys_platform == 'win32'
222226
uritemplate>=4.2.0
223227
urllib3>=2.6.3
224-
uvicorn>=0.47.0
228+
uvicorn>=0.49.0
225229
uvloop>=0.22.1
226-
virtualenv>=21.3.3
230+
virtualenv>=21.5.1
227231
wadler-lindig>=0.1.7
228232
websockets>=16.0
229233
werkzeug>=3.1.8
230234
wheel>=0.46.3
231-
wrapt>=2.1.2
232-
xxhash>=3.7.0
235+
wrapt>=2.2.2
236+
xxhash>=3.7.1
233237
yarl>=1.24.2
234238
zipp>=3.23.1
235239
zstandard>=0.25.0

src/maxtext/common/checkpointing.py

Lines changed: 7 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -443,17 +443,6 @@ def create_orbax_checkpoint_manager(
443443
logger=orbax_logger,
444444
)
445445

446-
# Use Colocated Python checkpointing optimization (Single Controller only).
447-
if enable_single_controller and colocated_python_checkpointing:
448-
max_logging.log("Registering colocated python array handler")
449-
checkpointing_impl = ocp.pathways.CheckpointingImpl.from_options(
450-
use_colocated_python=True,
451-
)
452-
ocp.pathways.register_type_handlers(
453-
use_single_replica_array_handler=enable_single_replica_ckpt_restoring,
454-
checkpointing_impl=checkpointing_impl,
455-
)
456-
457446
max_logging.log("Checkpoint manager created!")
458447
return manager
459448

@@ -500,6 +489,7 @@ def create_orbax_emergency_replicator_checkpoint_manager(
500489
local_checkpoint_dir: str,
501490
save_interval_steps: int,
502491
global_mesh: jax.sharding.Mesh,
492+
colocated_python_checkpointing: bool = False,
503493
):
504494
"""Returns an emergency replicator checkpoint manager."""
505495
flags.FLAGS.experimental_orbax_use_distributed_process_id = True
@@ -509,6 +499,7 @@ def create_orbax_emergency_replicator_checkpoint_manager(
509499
epath.Path(local_checkpoint_dir),
510500
options=emergency_replicator_checkpoint_manager.ReplicatorCheckpointManagerOptions(
511501
save_interval_steps=save_interval_steps,
502+
use_colocated_python=colocated_python_checkpointing,
512503
),
513504
global_mesh=global_mesh,
514505
)
@@ -838,9 +829,7 @@ def map_to_pspec(data):
838829
(EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager),
839830
):
840831
checkpoint_path = str(checkpoint_manager.directory / str(step) / "items")
841-
with handle_checkpoint_mismatch(
842-
"restore NNX checkpoint", checkpoint_path
843-
):
832+
with handle_checkpoint_mismatch("restore NNX checkpoint", checkpoint_path):
844833
restored_nnx = _load_linen_checkpoint_into_nnx(
845834
checkpoint_path,
846835
abstract_unboxed_pre_state,
@@ -876,9 +865,7 @@ def map_to_pspec(data):
876865
EmergencyReplicatorCheckpointManager,
877866
),
878867
):
879-
restored = checkpoint_manager.restore(
880-
step, args=Composite(state=checkpoint_args)
881-
).state
868+
restored = checkpoint_manager.restore(step, args=Composite(state=checkpoint_args)).state
882869
_assert_no_shaped_dtype_struct(restored)
883870
return (
884871
restored,
@@ -906,9 +893,7 @@ def map_to_pspec(data):
906893
# Case 3: Default/Fallback case.
907894
# This case acts as a wildcard ('_') and matches if none of the preceding cases were met.
908895
case _:
909-
restored = checkpoint_manager.restore(
910-
step, args=Composite(items=checkpoint_args)
911-
)
896+
restored = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args))
912897
_assert_no_shaped_dtype_struct(restored)
913898
return (restored, None)
914899

@@ -918,9 +903,7 @@ def map_to_pspec(data):
918903
else:
919904
params = abstract_unboxed_pre_state.params
920905

921-
with handle_checkpoint_mismatch(
922-
"load parameters", load_parameters_from_path
923-
):
906+
with handle_checkpoint_mismatch("load parameters", load_parameters_from_path):
924907
restored_params = load_params_from_path(
925908
load_parameters_from_path,
926909
params,
@@ -932,9 +915,7 @@ def map_to_pspec(data):
932915
return None, restored_params
933916
elif load_full_state_from_path != "":
934917
max_logging.log(f"Loading full state from path: {load_full_state_from_path}")
935-
with handle_checkpoint_mismatch(
936-
"load full state", load_full_state_from_path
937-
):
918+
with handle_checkpoint_mismatch("load full state", load_full_state_from_path):
938919
restored_state = _load_full_state_from_path(
939920
path=load_full_state_from_path,
940921
abstract_unboxed_pre_state=abstract_unboxed_pre_state,

src/maxtext/utils/max_utils.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -246,6 +246,17 @@ def maybe_initialize_jax_distributed_system(raw_keys):
246246
return
247247
if raw_keys["enable_single_controller"]:
248248
max_logging.log("Skipping jax distributed system since its not needed for single controller.")
249+
if raw_keys["enable_multi_tier_checkpointing"]:
250+
max_logging.log("Initializing multi-tier checkpointing for single controller...")
251+
initialize_multi_tier_checkpointing(
252+
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
253+
backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"],
254+
run_name=raw_keys["run_name"],
255+
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
256+
data_parallelism=raw_keys["mtc_data_parallelism"],
257+
num_slices=raw_keys["num_slices"],
258+
use_colocated_python=True,
259+
)
249260
return
250261
if jax.distributed.is_initialized():
251262
max_logging.log("Jax distributed system is already initialized.")
@@ -290,6 +301,7 @@ def maybe_initialize_jax_distributed_system(raw_keys):
290301
run_name=raw_keys["run_name"],
291302
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
292303
data_parallelism=raw_keys["mtc_data_parallelism"],
304+
num_slices=raw_keys["num_slices"],
293305
)
294306
max_logging.log("Jax distributed system initialized on TPUs for multi-tier checkpointing!")
295307
elif raw_keys["enable_checkpointing"] and raw_keys["compile_topology_num_slices"] == -1:

src/maxtext/utils/train_utils.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
import subprocess
1919
import jax
2020
import functools
21+
import orbax.checkpoint.pathways as ocp_pathways
2122
from functools import partial
2223

2324
from flax import nnx
@@ -55,6 +56,7 @@ def create_checkpoint_manager(config, mesh, init_state_fn):
5556
config.local_checkpoint_directory,
5657
config.local_checkpoint_period,
5758
mesh,
59+
config.colocated_python_checkpointing,
5860
)
5961
elif config.enable_emergency_checkpoint:
6062
abstract_state, _, _ = maxtext_utils.get_abstract_state(config, mesh, init_state_fn, is_training=True)
@@ -97,6 +99,17 @@ def create_checkpoint_manager(config, mesh, init_state_fn):
9799
config.checkpoint_todelete_full_path,
98100
)
99101

102+
# Use Colocated Python checkpointing dispatchers optimization (Single Controller only).
103+
if checkpoint_manager is not None and config.enable_single_controller and config.colocated_python_checkpointing:
104+
max_logging.log("Registering colocated python array handler")
105+
checkpointing_impl = ocp_pathways.CheckpointingImpl.from_options(
106+
use_colocated_python=True,
107+
)
108+
ocp_pathways.register_type_handlers(
109+
use_single_replica_array_handler=config.enable_single_replica_ckpt_restoring,
110+
checkpointing_impl=checkpointing_impl,
111+
)
112+
100113
return checkpoint_manager
101114

102115

0 commit comments

Comments
 (0)