Skip to content

Commit 436d7c4

Browse files
angel-coreGoogle-ML-Automation
authored andcommitted
Migrate Tunix's SFT Checkpoint Manager usage of Orbax V0 Checkpoint Manager to V1 Checkpointer.
Reverts f835ffb PiperOrigin-RevId: 949836450
1 parent e926028 commit 436d7c4

6 files changed

Lines changed: 223 additions & 76 deletions

File tree

docs/tutorials/posttraining/knowledge_distillation.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -234,7 +234,7 @@ python3 -m maxtext.checkpoint_conversion.to_maxtext \
234234
The online distillation trainer depends on Tunix. The XPK launcher script ([`scripts/run_distill_xpk.sh`](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh)) contains a `prep_image` step that layers Tunix on top of the MaxText base image. For local runs, install the same pin used by the launcher — the default `TUNIX_SOURCE` in `run_distill_xpk.sh` is the source of truth. As of this writing:
235235

236236
```bash
237-
pip install "git+https://github.com/google/tunix@110932a8395086511228483312131841521695c1"
237+
pip install "git+https://github.com/google/tunix@348959d18a4a09c75e58a7d49aec9d8b0eb4a8b6"
238238
```
239239

240240
> **Note:** The commit pin above will drift as the launcher is updated. Before installing, check the `TUNIX_SOURCE` default in [`run_distill_xpk.sh`](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh) and use that spec. Once a Tunix PyPI release ships, this will become a versioned `google-tunix==<ver>` install.
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,3 @@
1-
google-tunix @ https://github.com/google/tunix/archive/683256db1a0919b5cfd46cee52cebc96331494fb.zip
1+
google-tunix @ https://github.com/google/tunix/archive/348959d18a4a09c75e58a7d49aec9d8b0eb4a8b6.zip
22
tpu-inference @ https://github.com/vllm-project/tpu-inference/archive/4d08971683a64fd796a1b9fd0bb71128188882d5.zip
33
vllm @ git+https://github.com/vllm-project/vllm@2131b597b18d051dced4c4a605d362fa37f46ed1

src/maxtext/common/checkpointing.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515

1616
"""Create an Orbax CheckpointManager with specified (Async or not) Checkpointer."""
1717

18+
import asyncio
1819
import time
1920
from typing import Any, Optional
2021

@@ -169,6 +170,44 @@ class GrainCheckpointRestore(ocp.args.CheckpointArgs):
169170
process_count: Optional[int] = None
170171

171172

173+
class GrainCheckpointable(ocp_v1.StatefulCheckpointable):
174+
"""Adapts `GrainCheckpointHandler` to Orbax v1's `StatefulCheckpointable`."""
175+
176+
def __init__(
177+
self,
178+
*,
179+
save_args: GrainCheckpointSave | None = None,
180+
restore_args: GrainCheckpointRestore | None = None,
181+
):
182+
self._handler = GrainCheckpointHandler()
183+
self._save_args = save_args
184+
self._restore_args = restore_args
185+
186+
async def save(self, directory):
187+
"""Saves the Grain iterator state to the given directory."""
188+
# `GrainCheckpointHandler.save` snapshots iterator state (`get_state`) AND
189+
# writes it; both must happen in this (blocking) save phase, NOT in the
190+
# returned background coroutine.
191+
path = await directory.await_creation()
192+
self._handler.save(path, args=self._save_args)
193+
194+
async def _committed(): # nothing left for the background commit phase
195+
return None
196+
197+
return _committed()
198+
199+
async def load(self, directory: epath.Path):
200+
"""Loads the Grain iterator state from the given directory."""
201+
handler, args = self._handler, self._restore_args
202+
203+
# This will be ran to completion so asynchronous portion is just for
204+
# compatibility with Orbax v1 API.
205+
async def _background_load():
206+
await asyncio.to_thread(handler.restore, directory, args=args)
207+
208+
return _background_load()
209+
210+
172211
def _weight_mismatches(want, have, path=()):
173212
"""Returns `(path, problem)` for each weight in `want` that `have` didn't restore faithfully.
174213

src/maxtext/trainers/post_train/distillation/distillation_utils.py

Lines changed: 24 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -31,8 +31,7 @@
3131

3232
from maxtext.utils import max_logging
3333
from maxtext.utils import maxtext_utils
34-
# Reuse MaxText's native checkpointing logic.
35-
from maxtext.common.checkpointing import GrainCheckpointHandler, GrainCheckpointSave, GrainCheckpointRestore
34+
from maxtext.common import checkpointing
3635
from tunix.sft import checkpoint_manager as tunix_checkpoint_manager
3736
from tunix.sft import peft_trainer
3837

@@ -651,8 +650,9 @@ def create_labels(self, targets, targets_segmentation=None, **kwargs):
651650
class MaxTextCheckpointManager(tunix_checkpoint_manager.CheckpointManager):
652651
"""Custom CheckpointManager that uses MaxText's native handlers.
653652
654-
This manager extends Tunix to support saving/restoring the MaxText input pipeline
655-
(Grain) alongside the model and optimizer.
653+
Model and optimizer are delegated to Tunix's v1 ``Checkpointer`` unchanged.
654+
The Grain input pipeline is added as an extra ``"iter"`` checkpointable via
655+
``GrainCheckpointable``, which wraps MaxText's ``GrainCheckpointHandler``.
656656
"""
657657

658658
def __init__(
@@ -666,32 +666,6 @@ def __init__(
666666
self.student_config = student_config
667667
self._iterator = raw_iterator
668668

669-
# Re-initialize internal Orbax manager with MaxText's Grain handler
670-
# pylint: disable=access-member-before-definition
671-
# pytype: disable=attribute-error
672-
if self._checkpoint_manager is not None:
673-
root_directory = self._checkpoint_manager.directory
674-
675-
if options is None:
676-
options = getattr(self._checkpoint_manager, "options", None)
677-
678-
item_handlers = {
679-
"model_params": checkpoint.PyTreeCheckpointHandler(),
680-
"optimizer_state": checkpoint.PyTreeCheckpointHandler(),
681-
"custom_metadata": checkpoint.JsonCheckpointHandler(),
682-
# Use MaxText's handler for the iterator
683-
"iter": GrainCheckpointHandler(),
684-
}
685-
686-
self._checkpoint_manager.close()
687-
self._checkpoint_manager = checkpoint.CheckpointManager(
688-
root_directory,
689-
item_handlers=item_handlers,
690-
options=options,
691-
)
692-
# pytype: enable=attribute-error
693-
# pylint: enable=access-member-before-definition
694-
695669
def save(
696670
self,
697671
step,
@@ -701,10 +675,8 @@ def save(
701675
force=False,
702676
custom_metadata=None,
703677
):
704-
"""Saves the checkpoint including the input pipeline state (if available)."""
705-
if self._checkpoint_manager is None:
706-
return False
707-
if not force and not self._checkpoint_manager.should_save(step):
678+
"""Saves model, optimizer and the Grain input pipeline state."""
679+
if self._checkpointer is None:
708680
return False
709681

710682
# Standard Tunix Logic for Model/Optimizer.
@@ -715,21 +687,12 @@ def save(
715687
else:
716688
params = nnx.state(target_model)
717689

718-
# Define standard SaveArgs once to reuse
719-
default_save_args = checkpoint.SaveArgs()
720-
cp_save_args = {
721-
"model_params": checkpoint.args.PyTreeSave(
722-
item=params, save_args=jax.tree.map(lambda _: default_save_args, params)
723-
),
724-
}
725-
# Exclude optimizer state if the flag is set OR if learn_to_init_mode is active.
690+
checkpointables: dict[str, Any] = {"model_params": params}
691+
# Exclude optimizer state when learn_to_init_mode is active.
726692
exclude_opt = self.student_config.learn_to_init_mode
727693

728694
if optimizer is not None and not exclude_opt:
729-
optimizer_state = nnx.state(optimizer, nnx.optimizer.OptState)
730-
cp_save_args["optimizer_state"] = checkpoint.args.PyTreeSave(
731-
item=optimizer_state, save_args=jax.tree.map(lambda _: default_save_args, optimizer_state)
732-
)
695+
checkpointables["optimizer_state"] = nnx.state(optimizer, nnx.optimizer.OptState)
733696

734697
if self._iterator is not None:
735698
# Follow MaxText's logic to handle multi-process saving
@@ -747,15 +710,11 @@ def save(
747710
local_iter = data_iter.local_iterator if hasattr(data_iter, "local_iterator") else data_iter
748711
grain_iters_to_save.append((local_iter, process_index, process_count_total))
749712

750-
# Use GrainCheckpointSave wrapper
751-
cp_save_args["iter"] = GrainCheckpointSave(item=grain_iters_to_save) # pyrefly: ignore[bad-assignment]
713+
checkpointables["iter"] = checkpointing.GrainCheckpointable(
714+
save_args=checkpointing.GrainCheckpointSave(item=grain_iters_to_save) # pyrefly: ignore[bad-assignment]
715+
)
752716

753-
return self._checkpoint_manager.save(
754-
step,
755-
args=checkpoint.args.Composite(**cp_save_args),
756-
custom_metadata=custom_metadata or {},
757-
force=force,
758-
)
717+
return self._save_checkpointables(step, checkpointables, force, custom_metadata)
759718

760719
def maybe_restore( # pyrefly: ignore[bad-override]
761720
self,
@@ -770,12 +729,12 @@ def maybe_restore( # pyrefly: ignore[bad-override]
770729
Returns:
771730
(restored step, custom_metadata dict). Step is 0 if no checkpoint exists.
772731
"""
773-
if self._checkpoint_manager is None:
732+
if self._checkpointer is None:
774733
return 0, {}
775734

776735
target_model = getattr(model, "student_model", model)
777736

778-
step, _ = super().maybe_restore(
737+
step, custom_metadata = super().maybe_restore(
779738
model=target_model, # pyrefly: ignore[bad-argument-type]
780739
optimizer=optimizer,
781740
restore_only_lora_params=restore_only_lora_params,
@@ -785,20 +744,14 @@ def maybe_restore( # pyrefly: ignore[bad-override]
785744

786745
max_logging.log(f"Restored from checkpoint step {step}.")
787746

788-
metadata = self._checkpoint_manager.metadata(step)
789-
if metadata and hasattr(metadata, "custom_metadata") and metadata.custom_metadata is not None:
790-
custom_metadata = metadata.custom_metadata
791-
else:
792-
custom_metadata = {}
793-
794-
return step, dict(custom_metadata)
747+
return step, dict(custom_metadata or {})
795748

796749
def restore_iterator(self):
797750
"""Restores the iterator using MaxText's logic."""
798-
if self._checkpoint_manager is None or self._iterator is None:
751+
if self._checkpointer is None or self._iterator is None:
799752
return None
800753

801-
step = self._checkpoint_manager.latest_step()
754+
step = self.latest_step()
802755
if step is None:
803756
return None
804757

@@ -808,9 +761,10 @@ def restore_iterator(self):
808761
data_iter = self._iterator
809762
local_iter = data_iter.local_iterator if hasattr(data_iter, "local_iterator") else data_iter
810763

811-
restore_args = GrainCheckpointRestore(item=local_iter)
812-
813-
self._checkpoint_manager.restore(step, args=checkpoint.args.Composite(iter=restore_args))
764+
self._checkpointer.load_checkpointables(
765+
step,
766+
{"iter": checkpointing.GrainCheckpointable(restore_args=checkpointing.GrainCheckpointRestore(item=local_iter))},
767+
)
814768
# Since Grain restores in-place via set_state(), we return the original object
815769
return self._iterator
816770

@@ -820,5 +774,5 @@ def restore_iterator(self):
820774

821775
def wait_until_finished(self):
822776
"""Blocks until all outstanding checkpoint operations are complete."""
823-
if self._checkpoint_manager is not None:
824-
self._checkpoint_manager.wait_until_finished()
777+
if self._checkpointer is not None:
778+
self._checkpointer.wait()

src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,7 @@
105105
#
106106
# Image pinning (used by prep_image):
107107
# TUNIX_SOURCE pip-installable spec for tunix.
108-
# default: git+https://github.com/google/tunix@110932a8395086511228483312131841521695c1
108+
# default: git+https://github.com/google/tunix@348959d18a4a09c75e58a7d49aec9d8b0eb4a8b6
109109
# Use "google-tunix==<ver>" once a pypi release ships with the
110110
# multi-host shard_input fix.
111111
# JAX_PIN default: 0.10.0 — version to pin back after tunix deps resolve.
@@ -164,7 +164,7 @@ require_env() {
164164
: "${DISTILL_LAYER_INDICES:=[0,1,2,3,4,5,6,7]}"
165165

166166
# Image pinning (used by prep_image).
167-
: "${TUNIX_SOURCE:=git+https://github.com/google/tunix@110932a8395086511228483312131841521695c1}"
167+
: "${TUNIX_SOURCE:=git+https://github.com/google/tunix@348959d18a4a09c75e58a7d49aec9d8b0eb4a8b6}"
168168
: "${JAX_PIN:=0.10.0}"
169169
: "${JAXLIB_PIN:=0.10.0}"
170170
: "${LIBTPU_PIN:=0.0.39}"

0 commit comments

Comments
 (0)