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
698 changes: 322 additions & 376 deletions src/maxtext/common/checkpointing.py

Large diffs are not rendered by default.

321 changes: 16 additions & 305 deletions src/maxtext/common/grain_utility.py

Large diffs are not rendered by default.

4 changes: 3 additions & 1 deletion src/maxtext/configs/base.yml
Original file line number Diff line number Diff line change
Expand Up @@ -87,7 +87,9 @@ checkpoint_storage_use_zarr3: true
# default concurrent gb for PytreeCheckpointHandler is 96GB
checkpoint_storage_concurrent_gb: 96

# Bool flag for enabling Orbax v1.
# TODO: b/529622681 - Remove deprecated settings.
# DEPRECATED: Orbax v1 is now always used for checkpointing; this flag is
# ignored and will be removed in a future release.
enable_orbax_v1: false
# function for processing loaded checkpoint dict into a format maxtext can understand. (for other formats, i.e. safetensors)
checkpoint_conversion_fn: none
Expand Down
13 changes: 12 additions & 1 deletion src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
import jax
from maxtext.common.common_types import AttentionType, DecoderBlockType, ReorderStrategy, ShardMode, CustomRule, VisionEncoderBlockType
from maxtext.utils import gcs_utils
from maxtext.utils import max_logging
from maxtext.utils import max_utils
from maxtext.utils import elastic_utils
from maxtext.utils.globals import MAXTEXT_ASSETS_ROOT, HF_IDS
Expand Down Expand Up @@ -392,7 +393,10 @@ class Checkpointing(BaseModel):
description="Set to True if reading from a saved AQT quantized checkpoint.",
)
save_quantized_params_path: PathStr = Field("", description="Path to save params quantized on the fly.")
enable_orbax_v1: bool = Field(False, description="Bool flag for enabling Orbax v1.")
# TODO: b/529622681 - Remove deprecated settings.
enable_orbax_v1: bool = Field(
False, description="DEPRECATED: Orbax v1 is always used for checkpointing; this flag is ignored."
)
checkpoint_conversion_fn: None | str = Field(None, description="Function for processing loaded checkpoint dict.")
source_checkpoint_layout: Literal["orbax", "safetensors", "safetensors_dynamic"] = Field(
"orbax", description="The layout of the source checkpoint to load."
Expand Down Expand Up @@ -3589,6 +3593,13 @@ def get_num_target_devices():
"Please migrate to Qwix by setting use_qwix_quantization=True."
)

# Deprecated no-op: Orbax v1 is now the only checkpointing path.
if self.enable_orbax_v1:
max_logging.log(
"WARNING: enable_orbax_v1 is deprecated and ignored — Orbax v1 is now always used for "
"checkpointing. Remove the flag from your config; it will be deleted in a future release."
)

# Default quantization sharding count to number of local devices if not set.
if self.quantization_local_shard_count == -1:
try:
Expand Down
30 changes: 19 additions & 11 deletions src/maxtext/trainers/diloco/utils/spmd_diloco_checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,11 +19,12 @@
from flax import nnx
import jax
import jax.numpy as jnp
from maxtext.common import checkpoint_context
from maxtext.common import train_state_nnx
from maxtext.trainers.diloco import diloco
from maxtext.trainers.diloco.utils.nnx_state_utils import replace_nnx_model_params
import optax
import orbax.checkpoint as ocp
from orbax.checkpoint import v1 as ocp


# pylint: disable=too-many-positional-arguments
Expand All @@ -37,17 +38,24 @@ def restore_diloco_checkpoint(
) -> Any:
"""Restores a DiLoCo checkpoint into a DiLoCoTrainState."""
diloco_abstract = to_diloco_checkpoint_dict(abstract_nnx_state, config=config)
ckptr = ocp.Checkpointer(
ocp.PyTreeCheckpointHandler(
restore_concurrent_gb=checkpoint_storage_concurrent_gb,
save_concurrent_gb=checkpoint_storage_concurrent_gb,
use_ocdbt=use_ocdbt,
use_zarr3=use_zarr3,
)
# Orbax v1 refuses to read an item subdirectory directly (the step root carries the
# checkpoint indicator); normalize the documented ".../<step>/items" form to its root
# and load the checkpointable by name below. A v0-written flat pytree dir has no
# "items" child and is read directly.
root = epath.Path(str(path).rstrip("/"))
if root.name == "items":
root = root.parent
context = checkpoint_context.build_context(
use_ocdbt=use_ocdbt,
use_zarr3=use_zarr3,
checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb,
partial_load=True,
)
restore_args = ocp.checkpoint_utils.construct_restore_args(diloco_abstract)
restored = ocp.args.PyTreeRestore(item=diloco_abstract, restore_args=restore_args, partial_restore=True)
restored = ckptr.restore(epath.Path(path), args=restored)
with context:
checkpointable_name = "items" if (root / "items").exists() else None
restored = ocp.load(
root, diloco_abstract, checkpointable_name=checkpointable_name
) # pyrefly: ignore[bad-argument-type]
return from_diloco_checkpoint_dict(restored, abstract_nnx_state, config=config)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -652,7 +652,7 @@ class MaxTextCheckpointManager(tunix_checkpoint_manager.CheckpointManager):

Model and optimizer are delegated to Tunix's v1 ``Checkpointer`` unchanged.
The Grain input pipeline is added as an extra ``"iter"`` checkpointable via
``GrainCheckpointable``, which wraps MaxText's ``GrainCheckpointHandler``.
``GrainCheckpointable``, which implements Orbax's ``StatefulCheckpointable``.
"""

def __init__(
Expand Down Expand Up @@ -710,9 +710,7 @@ def save(
local_iter = data_iter.local_iterator if hasattr(data_iter, "local_iterator") else data_iter
grain_iters_to_save.append((local_iter, process_index, process_count_total))

checkpointables["iter"] = grain_utility.GrainCheckpointable(
save_args=grain_utility.GrainCheckpointSave(item=grain_iters_to_save) # pyrefly: ignore[bad-assignment]
)
checkpointables["iter"] = grain_utility.GrainCheckpointable(grain_iters_to_save)

return self._save_checkpointables(step, checkpointables, force, custom_metadata)

Expand Down Expand Up @@ -763,7 +761,7 @@ def restore_iterator(self):

self._checkpointer.load_checkpointables(
step,
{"iter": grain_utility.GrainCheckpointable(restore_args=grain_utility.GrainCheckpointRestore(item=local_iter))},
{"iter": grain_utility.GrainCheckpointable(local_iter)},
)
# Since Grain restores in-place via set_state(), we return the original object
return self._iterator
Expand Down
5 changes: 4 additions & 1 deletion src/maxtext/utils/train_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,9 @@ def create_training_optimizer(config, model):
def create_checkpoint_manager(config, mesh, init_state_fn):
"""Creates the init_rng, optimizer, learning rate schedule, and checkpoint manager."""
# pass in model for muon
logger = checkpointing.setup_checkpoint_logger(config)
# `setup_checkpoint_logger` only emits a deprecation warning now (Orbax v1 logs
# internally) and always returns None; we still pass it through for API parity.
logger = checkpointing.setup_checkpoint_logger(config) # pylint: disable=assignment-from-no-return
if config.enable_multi_tier_checkpointing:
checkpoint_manager = emergency_checkpointing.create_replicator_checkpoint_manager(
config.local_checkpoint_directory,
Expand Down Expand Up @@ -98,6 +100,7 @@ def create_checkpoint_manager(config, mesh, init_state_fn):
config.enable_autocheckpoint,
config.checkpoint_todelete_subdir,
config.checkpoint_todelete_full_path,
config.checkpoint_storage_target_data_file_size_bytes,
)

# Use Colocated Python checkpointing dispatchers optimization (Single Controller only).
Expand Down
4 changes: 2 additions & 2 deletions tests/integration/diloco_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -533,7 +533,7 @@ def test_diloco_checkpoint_saving_and_normal_resume(self):
use_zarr3=True,
)
checkpointing.save_checkpoint(mgr, 10, diloco_state, config, force=True)
mgr.wait_until_finished()
checkpointing.wait_until_finished(mgr)

items_path = os.path.join(temp_dir, "10", "items")

Expand Down Expand Up @@ -617,7 +617,7 @@ def test_diloco_automatic_checkpoint_resumption(self):
use_zarr3=True,
)
checkpointing.save_checkpoint(mgr, 5, diloco_state, config, force=True)
mgr.wait_until_finished()
checkpointing.wait_until_finished(mgr)

# Create new checkpoint manager for resumption (simulating next run with same run_name / checkpoint_dir)
resume_mgr = checkpointing.create_orbax_checkpoint_manager(
Expand Down
13 changes: 7 additions & 6 deletions tests/post_training/unit/lora_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,7 +332,7 @@ def test_sync_lora_metadata_default_syncs(self):
mock_metadata = mock.MagicMock()
mock_metadata.custom_metadata = {"lora": {"lora_rank": 32, "lora_alpha": 64.0}}

with mock.patch("orbax.checkpoint.StandardCheckpointer.metadata", return_value=mock_metadata):
with mock.patch.object(checkpointing.ocp, "checkpointables_metadata", return_value=mock_metadata):
lora_utils.sync_lora_metadata(cfg)
self.assertEqual(cfg.lora.lora_rank, 32)
self.assertEqual(cfg.lora.lora_alpha, 64.0)
Expand All @@ -350,7 +350,7 @@ def test_sync_lora_metadata_matching_passes(self):
mock_metadata = mock.MagicMock()
mock_metadata.custom_metadata = {"lora": {"lora_rank": 32, "lora_alpha": 64.0}}

with mock.patch("orbax.checkpoint.StandardCheckpointer.metadata", return_value=mock_metadata):
with mock.patch.object(checkpointing.ocp, "checkpointables_metadata", return_value=mock_metadata):
# Should not raise ValueError
lora_utils.sync_lora_metadata(cfg)
self.assertEqual(cfg.lora.lora_rank, 32)
Expand All @@ -369,7 +369,7 @@ def test_sync_lora_metadata_rank_mismatch_fails(self):
mock_metadata = mock.MagicMock()
mock_metadata.custom_metadata = {"lora": {"lora_rank": 32, "lora_alpha": 64.0}}

with mock.patch("orbax.checkpoint.StandardCheckpointer.metadata", return_value=mock_metadata):
with mock.patch.object(checkpointing.ocp, "checkpointables_metadata", return_value=mock_metadata):
with self.assertRaisesRegex(ValueError, "Configured lora_rank .* does not match"):
lora_utils.sync_lora_metadata(cfg)

Expand All @@ -386,7 +386,7 @@ def test_sync_lora_metadata_alpha_mismatch_fails(self):
mock_metadata = mock.MagicMock()
mock_metadata.custom_metadata = {"lora": {"lora_rank": 32, "lora_alpha": 64.0}}

with mock.patch("orbax.checkpoint.StandardCheckpointer.metadata", return_value=mock_metadata):
with mock.patch.object(checkpointing.ocp, "checkpointables_metadata", return_value=mock_metadata):
with self.assertRaisesRegex(ValueError, "Configured lora_alpha .* does not match"):
lora_utils.sync_lora_metadata(cfg)

Expand All @@ -398,11 +398,12 @@ def test_save_checkpoint_passes_metadata(self):
)
mock_manager = mock.MagicMock()
mock_state = mock.MagicMock()
mock_manager.use_async = False

with mock.patch("jax.block_until_ready"):
checkpointing.save_checkpoint(mock_manager, step=10, state=mock_state, config=cfg)
mock_manager.save.assert_called_once()
_, kwargs = mock_manager.save.call_args
mock_manager.save_checkpointables.assert_called_once()
_, kwargs = mock_manager.save_checkpointables.call_args
self.assertIn("custom_metadata", kwargs)
self.assertEqual(kwargs["custom_metadata"]["lora"], cfg.lora.model_dump())

Expand Down
71 changes: 48 additions & 23 deletions tests/unit/checkpointing_nnx_load_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,8 +240,8 @@ class TestManagerRestoreParity(unittest.TestCase):
"""

def _manager(self, restored):
manager = mock.MagicMock(spec=ocp.CheckpointManager)
manager.restore.return_value = restored
manager = mock.MagicMock(spec=checkpointing.ocp.training.Checkpointer)
manager.load_checkpointables.return_value = restored
return manager

def _linen_abstract(self):
Expand Down Expand Up @@ -273,7 +273,7 @@ def test_nnx_restores_the_linen_layout_and_returns_an_nnx_state(self):
restored, _ = self._load(manager, abstract)

# Going in: the manager is asked for the Linen on-disk layout, not the NNX one.
item = manager.restore.call_args.kwargs["args"]["items"].item
item = manager.load_checkpointables.call_args.args[1]["items"]
self.assertEqual(set(item) & {"params", "step"}, {"params", "step"})
self.assertNotIn("model", item)
# Coming out: reshaped back under `items`, the same key the Linen path returns.
Expand All @@ -287,34 +287,32 @@ def test_linen_restore_target_and_return_are_untouched(self):

restored, _ = self._load(manager, abstract)

self.assertIs(manager.restore.call_args.kwargs["args"]["items"].item, abstract)
self.assertIs(manager.load_checkpointables.call_args.args[1]["items"], abstract)
self.assertIs(restored, sentinel) # returned exactly as the manager gave it

def test_grain_case_converts_items_and_passes_the_iterator_through(self):
"""The grain branch is shared too: NNX only reshapes `items`, leaving the iterator element alone."""
"""The grain branch is shared too: NNX only reshapes `items`; the iterator restores in place."""
abstract = _abstract_nnx_state()
saved = {"params": {"params": {"linear": {"kernel": jnp.ones((2, 1)), "bias": jnp.array([5.0])}}}}
manager = self._manager(None)
manager = self._manager({"items": saved, "iter": mock.Mock()})

with mock.patch.object(
checkpointing.grain_utility, "restore_grain_iterator", return_value=({"items": saved}, None)
) as m:
with mock.patch.object(checkpointing.grain_utility, "for_restore", return_value=mock.Mock()) as m:
restored, iterator = self._load(manager, abstract, dataset_type="grain", data_iterator=mock.MagicMock())

m.assert_called_once()
# The iterator checkpointable was requested alongside the state.
self.assertIn("iter", manager.load_checkpointables.call_args.args[1])
self.assertIsNone(iterator)
self.assertIsInstance(restored["items"], nnx.State)
self.assertTrue(jnp.array_equal(restored["items"].to_pure_dict()["model"]["linear"]["bias"], jnp.array([5.0])))

def test_grain_case_is_untouched_for_linen(self):
abstract = self._linen_abstract()
sentinel = ({"items": {"params": abstract.params}}, None)
manager = self._manager(None)
sentinel = {"items": {"params": abstract.params}}
manager = self._manager(sentinel)

with mock.patch.object(checkpointing.grain_utility, "restore_grain_iterator", return_value=sentinel):
with mock.patch.object(checkpointing.grain_utility, "for_restore", return_value=mock.Mock()):
restored, iterator = self._load(manager, abstract, dataset_type="grain", data_iterator=mock.MagicMock())

self.assertIs(restored, sentinel[0])
self.assertIs(restored, sentinel)
self.assertIsNone(iterator)

def test_missing_weight_raises_on_the_standard_path(self):
Expand All @@ -336,10 +334,9 @@ def test_missing_weight_raises_on_the_emergency_path(self):
self.assertIn("linear/bias", str(ctx.exception))

def test_no_step_in_manager_falls_through_to_the_load_paths(self):
"""An empty manager (latest_step() is None) must not restore -- it falls through, for both types."""
manager = mock.MagicMock(spec=ocp.CheckpointManager)
manager.latest_step.return_value = None

"""An empty manager (latest is None) must not restore -- it falls through, for both types."""
manager = mock.MagicMock(spec=checkpointing.ocp.training.Checkpointer)
manager.latest = None
for abstract in (_abstract_nnx_state(), self._linen_abstract()):
restored, params = checkpointing.load_state_if_possible(
checkpoint_manager=manager,
Expand All @@ -351,7 +348,7 @@ def test_no_step_in_manager_falls_through_to_the_load_paths(self):
)
self.assertIsNone(restored)
self.assertIsNone(params)
manager.restore.assert_not_called()
manager.load_checkpointables.assert_not_called()


class TestResolveConversionFn(unittest.TestCase):
Expand Down Expand Up @@ -444,15 +441,16 @@ class TestSafetensorsFullStateIntoNNX(unittest.TestCase):
def _load(self, converted, abstract):
"""Runs the v1 safetensors branch, stubbing the read so only the conversion is under test."""
with (
mock.patch.object(checkpointing, "ocp_v1") as v1,
mock.patch.object(checkpointing, "ocp") as v1,
mock.patch.object(checkpointing, "sharding_utils") as shardings,
mock.patch.object(checkpointing, "checkpoint_context") as context,
):
v1.pytree_metadata.return_value = mock.Mock(metadata={"w": jax.ShapeDtypeStruct((1,), jnp.float32)})
v1.metadata.return_value = mock.Mock(metadata={"w": jax.ShapeDtypeStruct((1,), jnp.float32)})
shardings.construct_maximal_shardings.return_value = {"w": None}
context.build_context.return_value = mock.MagicMock() # a with-able context
return checkpointing._load_full_state_from_path( # pylint: disable=protected-access
path="gs://does-not-exist/hf",
abstract_unboxed_pre_state=abstract,
enable_orbax_v1=True,
checkpoint_conversion_fn=lambda _: converted,
source_checkpoint_layout="safetensors",
checkpoint_storage_concurrent_gb=8,
Expand Down Expand Up @@ -711,6 +709,33 @@ def test_weight_mismatches_finds_absent_sds_and_shape(self):
self.assertIn("missing", problems["a/b"])
self.assertIn("shape", problems["a/c"])

def test_weight_mismatches_ignores_missing_when_check_missing_is_false(self):
want = {
"a": {
"k": jax.ShapeDtypeStruct((2,), jnp.float32),
"b": jax.ShapeDtypeStruct((1,), jnp.float32),
"c": jax.ShapeDtypeStruct((3,), jnp.float32),
}
}
# b absent, c wrong shape
have = {"a": {"k": jnp.ones((2,)), "c": jnp.ones((4,))}}
problems = dict(checkpointing._weight_mismatches(want, have, check_missing=False)) # pylint: disable=protected-access
self.assertEqual(list(problems.keys()), ["a/c"])
self.assertIn("shape", problems["a/c"])

def test_weight_mismatches_detects_structural_mismatch(self):
want = {
"a": {
"k": jax.ShapeDtypeStruct((2,), jnp.float32),
}
}
# k is a dictionary instead of a tensor
have = {"a": {"k": {"nested_dict_instead": 1}}}
problems = dict(checkpointing._weight_mismatches(want, have)) # pylint: disable=protected-access
self.assertEqual(list(problems.keys()), ["a/k"])
self.assertIn("structural mismatch", problems["a/k"])
self.assertIn("dict", problems["a/k"])
Comment on lines +735 to +737

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟢 Low - Adding a corresponding test case for the inverse structural mismatch (where the model expects a dictionary, but the checkpoint provides a tensor) to ensure complete and robust test coverage for mismatch scenarios.
Suggested change
self.assertEqual(list(problems.keys()), ["a/k"])
self.assertIn("structural mismatch", problems["a/k"])
self.assertIn("dict", problems["a/k"])
self.assertEqual(list(problems.keys()), ["a/k"])
self.assertIn("structural mismatch", problems["a/k"])
self.assertIn("dict", problems["a/k"])
def test_weight_mismatches_detects_structural_mismatch_inverse(self):
want = {
"a": {
"k": {
"w": jax.ShapeDtypeStruct((2,), jnp.float32),
}
}
}
# k is a tensor instead of a dictionary
have = {"a": {"k": jnp.ones((2,))}}
problems = dict(checkpointing._weight_mismatches(want, have)) # pylint: disable=protected-access
self.assertEqual(list(problems.keys()), ["a/k"])
self.assertIn("structural mismatch", problems["a/k"])
self.assertIn("model expects a dict", problems["a/k"])


def test_expected_and_restored_params_splits_by_param_type(self):
"""Only nnx.Param weights land in `want`; rngs/dropout (nnx.RngState) are excluded from the check."""
model = _ModelDropout(nnx.Rngs(0))
Expand Down
Loading
Loading