diff --git a/src/maxtext/common/checkpointing.py b/src/maxtext/common/checkpointing.py index 7191190f2b..d46d8b89be 100644 --- a/src/maxtext/common/checkpointing.py +++ b/src/maxtext/common/checkpointing.py @@ -1,4 +1,4 @@ -# Copyright 2023–2025 Google LLC +# Copyright 2023–2026 Google LLC # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -16,7 +16,6 @@ """Create an Orbax CheckpointManager with specified (Async or not) Checkpointer.""" import contextlib -import datetime import importlib import os import time @@ -28,71 +27,71 @@ from flax.training import train_state -from grain.experimental import ElasticIterator import jax -from maxtext.checkpoint_conversion.utils.load_dynamic import load_safetensors_dynamic_state +from jax.experimental import multihost_utils +from maxtext.checkpoint_conversion.utils import load_dynamic +from maxtext.common import checkpoint_context from maxtext.common import emergency_checkpointing from maxtext.common import grain_utility from maxtext.common import train_state_nnx -from maxtext.input_pipeline.multihost_dataloading import MultiHostDataLoadIterator -from maxtext.input_pipeline.multihost_dataloading import RemoteIteratorWrapper -from maxtext.input_pipeline.synthetic_data_processing import PlaceHolderDataIterator +from maxtext.input_pipeline import multihost_dataloading +from maxtext.input_pipeline import synthetic_data_processing from maxtext.trainers.diloco.utils import spmd_diloco_checkpointing as diloco_checkpoint_utils from maxtext.utils import elastic_utils from maxtext.utils import exceptions from maxtext.utils import gcs_utils +from maxtext.utils import globals as maxtext_globals from maxtext.utils import max_logging -from maxtext.utils.globals import DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE -import orbax.checkpoint as ocp -from orbax.checkpoint import v1 as ocp_v1 +from orbax.checkpoint import v1 as ocp from orbax.checkpoint._src.arrays import sharding as sharding_utils -from orbax.checkpoint._src.checkpoint_managers import preservation_policy as preservation_policy_lib -from orbax.checkpoint._src.checkpoint_managers import save_decision_policy as save_decision_policy_lib -CheckpointManagerOptions = ocp.CheckpointManagerOptions -Composite = ocp.args.Composite -PyTreeCheckpointHandler = ocp.PyTreeCheckpointHandler +load_safetensors_dynamic_state = load_dynamic.load_safetensors_dynamic_state +PlaceHolderDataIterator = synthetic_data_processing.PlaceHolderDataIterator +MultiHostDataLoadIterator = multihost_dataloading.MultiHostDataLoadIterator +DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE = maxtext_globals.DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE + # Backward compatibility aliases for v0 emergency managers. EmergencyCheckpointManager = emergency_checkpointing.CheckpointManager EmergencyReplicatorCheckpointManager = emergency_checkpointing.ReplicatorCheckpointManager create_orbax_emergency_checkpoint_manager = emergency_checkpointing.create_emergency_checkpoint_manager create_orbax_emergency_replicator_checkpoint_manager = emergency_checkpointing.create_replicator_checkpoint_manager -# Union of CheckpointManager / the emergency factories return; used in type hints. -CheckpointManager = ocp.CheckpointManager | EmergencyCheckpointManager | EmergencyReplicatorCheckpointManager +# Union of v1 Checkpointer / the emergency factories return; used in type hints. +CheckpointManager = ocp.training.Checkpointer | EmergencyCheckpointManager | EmergencyReplicatorCheckpointManager -def _weight_mismatches(want, have, path=(), is_quantized_param=False): - """Returns `(path, problem)` for each weight in `want` that `have` didn't restore faithfully. +def _weight_mismatches(want, have, path=(), check_missing: bool = True, is_quantized_param: bool = False): + """Returns `(path, problem)` for each weight in `want` that `have` didn't restore. - A weight is wrong if the checkpoint didn't carry it -- absent, or left by Orbax as an - unmaterialized ShapeDtypeStruct -- or carried it at a different shape. Only the shape can - disagree: Orbax casts a restored array to the target's dtype. + If check_missing is True, this reports absent weights (post-load check). If check_missing is + False, missing weights are ignored and only shapes of matching keys are checked (the pre-load + check against checkpoint metadata, whose leaves carry shapes but no values). LoRA adapter + weights, rng streams, and quantized-param subtrees are allowed to be absent from the + checkpoint. Only shapes and structure can disagree: Orbax casts restored array dtypes. """ if isinstance(want, dict): out = [] is_quant = is_quantized_param or any(k in want for k in ("qvalue", "qarray")) for k, v in want.items(): - out.extend( - _weight_mismatches( - v, - have.get(k) if isinstance(have, dict) else None, - path + (k,), - is_quant, - ) - ) + nested = have.get(k) if isinstance(have, dict) else None + if check_missing or nested is not None: + out.extend(_weight_mismatches(v, nested, path + (k,), check_missing=check_missing, is_quantized_param=is_quant)) return out name = "/".join(str(p) for p in path) - if "lora_a" in name or "lora_b" in name or "rngs" in path or "rng" in path: - if have is None or isinstance(have, jax.ShapeDtypeStruct): - return [] - - if (have is None or isinstance(have, jax.ShapeDtypeStruct)) and is_quantized_param: + is_missing = have is None or (check_missing and isinstance(have, jax.ShapeDtypeStruct)) + if is_missing and ("lora_a" in name or "lora_b" in name or "rngs" in path or "rng" in path): return [] - - if have is None or isinstance(have, jax.ShapeDtypeStruct): - return [(name, f"missing (model expects {getattr(want, 'shape', '?')} {getattr(want, 'dtype', '?')})")] + if is_missing and is_quantized_param: + return [] + if is_missing: + return ( + [(name, f"missing (model expects {getattr(want, 'shape', '?')} {getattr(want, 'dtype', '?')})")] + if check_missing + else [] + ) want_shape, got_shape = getattr(want, "shape", None), getattr(have, "shape", None) + if want_shape is not None and got_shape is None: + return [(name, f"structural mismatch: model expects a tensor but checkpoint provides {type(have).__name__}")] if want_shape is not None and got_shape is not None and tuple(want_shape) != tuple(got_shape): return [(name, f"shape {tuple(got_shape)} but the model expects {tuple(want_shape)}")] return [] @@ -109,6 +108,19 @@ def _expected_and_restored_params(abstract_nnx_state, restored_linen): return want, have +def _raise_weight_problems(problems): + """Raises a ValueError naming each mismatched weight; returns if there are none.""" + # Ignore the weight mismatches in the custom projector so it can stay randomly initialized + if not problems: + return + lines = "\n".join(f" - '{p}': {why}" for p, why in problems) + raise ValueError( + "Checkpoint does not match the model:\n" + f"{lines}\n" + "Verify the checkpoint matches the model architecture (emb_dim, mlp_dim, num layers, scan_layers)." + ) + + def _is_custom_projector_problem(path: str, want: dict) -> bool: """Returns True if a weight mismatch belongs to a newly attached custom vision projector.""" parts = [p for p in path.replace(".", "/").split("/") if p and p != "params"] @@ -138,16 +150,8 @@ def _raise_on_weight_mismatch(want, have, config=None): want = _filter_lora_trainable_state(want) problems = _weight_mismatches(want, have) - # Ignore the weight mismatches in the custom projector so it can stay randomly initialized problems = [(p, why) for p, why in problems if not _is_custom_projector_problem(p, want)] - if not problems: - return - lines = "\n".join(f" - '{p}': {why}" for p, why in problems) - raise ValueError( - "Checkpoint does not match the model:\n" - f"{lines}\n" - "Verify the checkpoint matches the model architecture (emb_dim, mlp_dim, num layers, scan_layers)." - ) + _raise_weight_problems(problems) def _linen_items_to_nnx(restored_linen, abstract_nnx_state): @@ -179,12 +183,13 @@ def _load_linen_checkpoint_into_nnx( checkpoint_storage_concurrent_gb, use_ocdbt, use_zarr3, + enable_single_replica_ckpt_restoring: bool = False, config=None, ): """Restores a Linen-layout checkpoint into an NNX state (pure_nnx resume). Restores a Linen-shape target that includes `nnx_aux`, then reshapes back via - `_linen_items_to_nnx`. rngs/dropout/batch stats come from `items/nnx_aux` when + `_restored_linen_to_nnx`. rngs/dropout/batch stats come from `items/nnx_aux` when present, else keep their fresh init value. A genuinely-missing weight raises. """ max_logging.log(f"Restoring Linen-layout checkpoint into NNX state at {path}") @@ -201,17 +206,23 @@ def _load_linen_checkpoint_into_nnx( linen_abstract = train_state_nnx.to_checkpoint_dict(abstract_nnx_state) if config and getattr(getattr(config, "lora", None), "enable_lora", False): linen_abstract = _filter_lora_trainable_state(linen_abstract) - 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, - ) + context = checkpoint_context.build_context( + use_ocdbt=use_ocdbt, + use_zarr3=use_zarr3, + checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, + partial_load=True, + enable_single_replica_ckpt_restoring=enable_single_replica_ckpt_restoring, ) - restore_args = ocp.checkpoint_utils.construct_restore_args(linen_abstract) - restored = ocp.args.PyTreeRestore(item=linen_abstract, restore_args=restore_args, partial_restore=True) - restored = ckptr.restore(epath.Path(path), args=restored) + # Orbax v1 refuses to read an item subdirectory directly (the step root carries the + # checkpoint indicator); normalize the documented "...//items" form to its root + # and load the checkpointable by name. A v0-written flat pytree dir has no "items" + # child and is read directly. + root = epath.Path(_normalize_checkpoint_root(path)) + with context: + checkpointable_name = "items" if (root / "items").exists() else None + restored = ocp.load( + root, linen_abstract, checkpointable_name=checkpointable_name + ) # pyrefly: ignore[bad-argument-type] return _restored_linen_to_nnx(restored, abstract_nnx_state, config=config) @@ -269,12 +280,12 @@ def _resolve_conversion_fn(checkpoint_conversion_fn): def _load_full_state_from_path( path, abstract_unboxed_pre_state, - enable_orbax_v1, checkpoint_conversion_fn, source_checkpoint_layout, checkpoint_storage_concurrent_gb, use_ocdbt, use_zarr3, + enable_single_replica_ckpt_restoring: bool = False, maxtext_config=None, ): """Load full state from checkpoint at specified path. @@ -283,7 +294,6 @@ def _load_full_state_from_path( path: path to checkpoint abstract_unboxed_pre_state: an abstract state that Orbax matches type against. - enable_orbax_v1: whether to use orbax v1 or the previously supported v0. checkpoint_conversion_fn: user-provided function to convert checkpoint to maxtext-supported state. source_checkpoint_layout: String representation of the checkpoint layout of @@ -291,51 +301,15 @@ def _load_full_state_from_path( checkpoint_storage_concurrent_gb: concurrent GB for checkpoint byte I/O. use_ocdbt: Whether to use OCDBT format. use_zarr3: Whether to use Zarr3 format. + enable_single_replica_ckpt_restoring: bool flag for restoring checkpoint + with load-and-broadcast (single replica). Supported for Orbax format only. maxtext_config: Optional configuration dictionary/object. Returns: The loaded state. """ - - if enable_orbax_v1: - if source_checkpoint_layout == "orbax": - # pure_nnx saves in the Linen on-disk layout; reshape it back into the NNX state. - if isinstance(abstract_unboxed_pre_state, nnx.State): - return _load_linen_checkpoint_into_nnx( - path, - abstract_unboxed_pre_state, - checkpoint_storage_concurrent_gb, - use_ocdbt, - use_zarr3, - config=maxtext_config, - ) - context = ocp_v1.Context(checkpoint_layout=ocp_v1.options.CheckpointLayout.ORBAX) - with context: - return ocp_v1.load_pytree(path, abstract_unboxed_pre_state) - elif source_checkpoint_layout == "safetensors": - # Resolved first, so a bad config fails before the weights are read. - conversion_fn = _resolve_conversion_fn(checkpoint_conversion_fn) - context = ocp_v1.Context(checkpoint_layout=ocp_v1.options.CheckpointLayout.SAFETENSORS) - with context: - metadata = ocp_v1.pytree_metadata(path) - simple_abstract_state = metadata.metadata - shardings = sharding_utils.construct_maximal_shardings(simple_abstract_state) - - def combine_sharding(sds, shardings): - return jax.ShapeDtypeStruct(shape=sds.shape, dtype=sds.dtype, sharding=shardings) - - sharded_abstract_state = jax.tree.map(combine_sharding, simple_abstract_state, shardings) - pre_transformed_state = ocp_v1.load_pytree(path, sharded_abstract_state) - state = conversion_fn(pre_transformed_state) - # The conversion fn returns MaxText's on-disk (Linen) layout, which is what pure_nnx reads, - # so NNX needs the same reshape as every other restore. An NNX state passes through. - if isinstance(abstract_unboxed_pre_state, nnx.State) and not isinstance(state, nnx.State): - state = _restored_linen_to_nnx(state, abstract_unboxed_pre_state, config=maxtext_config) - return state - else: - raise ocp_v1.errors.InvalidLayoutError(f"Unknown checkpoint layout: {source_checkpoint_layout}") - else: - # pure_nnx saves in the Linen on-disk layout; reshape it back into the NNX state. + if source_checkpoint_layout == "orbax": + # pure_nnx checkpoints are stored in the Linen on-disk layout; reshape to NNX. if isinstance(abstract_unboxed_pre_state, nnx.State): return _load_linen_checkpoint_into_nnx( path, @@ -343,25 +317,46 @@ def combine_sharding(sds, shardings): checkpoint_storage_concurrent_gb, use_ocdbt, use_zarr3, + enable_single_replica_ckpt_restoring=enable_single_replica_ckpt_restoring, config=maxtext_config, ) - - # Original v0 logic. - p = epath.Path(path) - handler = ocp.PyTreeCheckpointHandler( - restore_concurrent_gb=checkpoint_storage_concurrent_gb, - save_concurrent_gb=checkpoint_storage_concurrent_gb, + context = checkpoint_context.build_context( use_ocdbt=use_ocdbt, use_zarr3=use_zarr3, + checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, + checkpoint_layout=ocp.options.CheckpointLayout.ORBAX, + enable_single_replica_ckpt_restoring=enable_single_replica_ckpt_restoring, ) - # Only Linen TrainState reaches here; nnx.State returned above. - restore_target = abstract_unboxed_pre_state - # Provide sharding info to ensure restoration returns JAX arrays (not NumPy arrays). - restore_args = jax.tree_util.tree_map( - lambda x: ocp.type_handlers.ArrayRestoreArgs(sharding=x.sharding), - restore_target, + with context: + return ocp.load(path, abstract_unboxed_pre_state) + + if source_checkpoint_layout == "safetensors": + if enable_single_replica_ckpt_restoring: + max_logging.warning("enable_single_replica_ckpt_restoring is not supported for safetensors layout.") + # Resolved first, so a bad config fails before the weights are read. + conversion_fn = _resolve_conversion_fn(checkpoint_conversion_fn) + context = checkpoint_context.build_context( + checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, + checkpoint_layout=ocp.options.CheckpointLayout.SAFETENSORS, ) - return ocp.Checkpointer(handler).restore(p, restore_target, restore_args=restore_args) + with context: + metadata = ocp.metadata(path) + simple_abstract_state = metadata.metadata + shardings = sharding_utils.construct_maximal_shardings(simple_abstract_state) + + def combine_sharding(sds, shardings): + return jax.ShapeDtypeStruct(shape=sds.shape, dtype=sds.dtype, sharding=shardings) + + sharded_abstract_state = jax.tree.map(combine_sharding, simple_abstract_state, shardings) + pre_transformed_state = ocp.load(path, sharded_abstract_state) + state = conversion_fn(pre_transformed_state) + # The conversion fn returns MaxText's on-disk (Linen) layout, which is what pure_nnx reads, + # so NNX needs the same reshape as every other restore. An NNX state passes through. + if isinstance(abstract_unboxed_pre_state, nnx.State) and not isinstance(state, nnx.State): + state = _restored_linen_to_nnx(state, abstract_unboxed_pre_state, config=maxtext_config) + return state + + raise ocp.errors.InvalidLayoutError(f"Unknown checkpoint layout: {source_checkpoint_layout}") def create_orbax_checkpoint_manager( @@ -379,68 +374,59 @@ def create_orbax_checkpoint_manager( enable_autocheckpoint: bool = False, todelete_subdir: str | None = None, todelete_full_path: str | None = None, + ocdbt_target_data_file_size_bytes: int | None = None, ): - """Returns specified Orbax (async or not) CheckpointManager or None if checkpointing is disabled.""" + """Returns an Orbax v1 training ``Checkpointer``, or None if checkpointing is disabled.""" if not enable_checkpointing: max_logging.log("Checkpointing disabled, not creating checkpoint manager.") return None - max_logging.log(f"Creating checkpoint manager with ocdbt={use_ocdbt} and zarr3={use_zarr3}") - - # Base configuration for all dataset types - item_names = ("items",) - # we need to use ocdbt and zarr3 to control max file size in the checkpoint - item_handlers = { - "items": PyTreeCheckpointHandler( - restore_concurrent_gb=checkpoint_storage_concurrent_gb, - save_concurrent_gb=checkpoint_storage_concurrent_gb, - use_ocdbt=use_ocdbt, - use_zarr3=use_zarr3, - ) - } - - if dataset_type is not None and dataset_type == "grain": - item_names += ("iter",) - item_handlers["iter"] = grain_utility.GrainCheckpointHandler() # pyrefly: ignore[bad-assignment] - - # local storage checkpoint needs parent directory created - p = gcs_utils.mkdir_and_check_permissions(checkpoint_dir) - if enable_continuous_checkpointing: - max_logging.log("Enabling policy for continuous checkpointing.") - save_decision_policy = save_decision_policy_lib.ContinuousCheckpointingPolicy() - elif enable_autocheckpoint: - max_logging.log("Enabling policy for autocheckpoint.") - save_decision_policy = save_decision_policy_lib.AnySavePolicy( - [ - save_decision_policy_lib.PreemptionCheckpointingPolicy(), - save_decision_policy_lib.FixedIntervalPolicy(save_interval_steps), - ] + # TODO: b/529622681 - Remove deprecated settings. + if orbax_logger is not None: + max_logging.warning( + "Cloud logging (enable_checkpoint_cloud_logger) is disabled because" + " Orbax v1 now configures its own logger internally. This config" + " setting is ignored and will be removed." ) - else: - max_logging.log("Enabling policy for fixed interval checkpointing.") - save_decision_policy = save_decision_policy_lib.FixedIntervalPolicy(interval=save_interval_steps) - preservation_policy = preservation_policy_lib.LatestN(max_num_checkpoints_to_keep) - - async_options = None - if enable_continuous_checkpointing: - async_options = ocp.AsyncOptions( - timeout_secs=int(datetime.timedelta(minutes=60).total_seconds()), + if dataset_type is not None: + max_logging.warning( + "Specifying dataset_type upon checkpointer creation is deprecated and" + " will be removed soon, this is now handled dynamically by Orbax" + " Checkpointer." ) - manager = ocp.CheckpointManager( - p, - item_names=item_names, - item_handlers=item_handlers, - options=CheckpointManagerOptions( - create=True, - enable_async_checkpointing=use_async, - save_decision_policy=save_decision_policy, - preservation_policy=preservation_policy, - async_options=async_options, - todelete_subdir=todelete_subdir, - todelete_full_path=todelete_full_path, + + max_logging.log(f"Creating checkpointer with ocdbt={use_ocdbt} and zarr3={use_zarr3}") + + validated_path = gcs_utils.mkdir_and_check_permissions(checkpoint_dir) + + if ocdbt_target_data_file_size_bytes is None: + ocdbt_target_data_file_size_bytes = DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE + + context = checkpoint_context.build_context( + use_ocdbt=use_ocdbt, + use_zarr3=use_zarr3, + ocdbt_target_data_file_size_bytes=ocdbt_target_data_file_size_bytes, + checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, + enable_continuous_checkpointing=enable_continuous_checkpointing, + todelete_full_path=todelete_full_path, + todelete_subdir=todelete_subdir, + partial_load=True, + ) + + manager = ocp.training.Checkpointer( + validated_path, + context=context, + save_decision_policy=checkpoint_context.build_save_decision_policy( + save_interval_steps=save_interval_steps, + enable_continuous_checkpointing=enable_continuous_checkpointing, + enable_autocheckpoint=enable_autocheckpoint, + ), + preservation_policy=checkpoint_context.build_preservation_policy( + max_to_keep=max_num_checkpoints_to_keep, ), - logger=orbax_logger, ) + # Necessary bridge to support v0 backward compatibility. + manager.use_async = use_async # pyrefly: ignore[missing-attribute] max_logging.log("Checkpoint manager created!") return manager @@ -455,17 +441,42 @@ def print_save_message(step, async_checkpointing): def latest_step(checkpoint_manager): """Latest saved step or None, across the v0 emergency manager and the v1 Checkpointer.""" - return checkpoint_manager.latest_step() + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + return checkpoint_manager.latest_step() + else: + latest = checkpoint_manager.latest + return latest.step if latest is not None else None + + +def all_steps(checkpoint_manager): + """All saved steps, across the v0 emergency manager and the v1 Checkpointer.""" + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + return checkpoint_manager.all_steps() + return [checkpoint.step for checkpoint in checkpoint_manager.checkpoints] def wait_until_finished(checkpoint_manager): """Blocks until pending saves finish, across the v0 emergency manager and the v1 Checkpointer.""" - return checkpoint_manager.wait_until_finished() + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + checkpoint_manager.wait_until_finished() + else: + checkpoint_manager.wait() def reached_preemption(checkpoint_manager, step: int) -> bool: - """Whether a preemption sync point has been reached at `step`, across the v0 emergency manager and the v1 Checkpointer.""" - return checkpoint_manager.reached_preemption(step) + """Whether a preemption sync point has been reached at ``step`` (manager-agnostic).""" + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + return checkpoint_manager.reached_preemption(step) + else: + return multihost_utils.reached_preemption_sync_point(step) + + +def _normalize_checkpoint_root(path_str): + """Lifts a v0-convention pytree path ("...//items") to its checkpoint root.""" + path_str = str(path_str).rstrip("/") + if path_str == "items": + return "." + return path_str.removesuffix("/items") def load_state_if_possible( @@ -497,13 +508,13 @@ def load_state_if_possible( manager, load full state from a full state checkpoint at this path. abstract_unboxed_pre_state: an unboxed, abstract TrainState that Orbax matches type against. - enable_single_replica_ckpt_restoring: bool flag for restoring checkpoitn - with SingleReplicaArrayHandler + enable_single_replica_ckpt_restoring: bool flag for restoring checkpoint + with load-and-broadcast (single replica). Supported for Orbax format only. checkpoint_storage_concurrent_gb: concurrent GB for checkpoint byte I/O. enable_orbax_v1: bool flag for enabling Orbax v1. checkpoint_conversion_fn: function for converting checkpoint to Orbax v1. source_checkpoint_layout: Optional checkpoint context to use for loading, - provided in string format with the default being "orbax". + provided in string format with the default being "orbax". Returns: A tuple of (train_state, train_state_params) where full_train_state captures @@ -512,6 +523,16 @@ def load_state_if_possible( set. """ + # TODO: b/529622681 - Remove deprecated settings. + if enable_orbax_v1: + max_logging.warning( + "enable_orbax_v1 is deprecated and will be removed, as Orbax v1 is now the default checkpointing API." + ) + + if load_parameters_from_path: + load_parameters_from_path = _normalize_checkpoint_root(load_parameters_from_path) + if load_full_state_from_path: + load_full_state_from_path = _normalize_checkpoint_root(load_full_state_from_path) # pure_nnx saves in the Linen on-disk layout, so every branch below loads the same tree Linen # does: the NNX abstract is converted to that layout going in, and what comes back is reshaped # into the NNX state on the way out. @@ -524,30 +545,6 @@ def load_state_if_possible( if step is not None: max_logging.log(f"restoring from this run's directory step {step}") - def map_to_pspec(data): - if not enable_single_replica_ckpt_restoring: - return ocp.type_handlers.ArrayRestoreArgs(sharding=data.sharding) - pspec = data.sharding.spec - mesh = data.sharding.mesh - replica_axis_index = 0 - replica_devices = grain_utility.replica_devices(mesh.devices, replica_axis_index) - replica_mesh = jax.sharding.Mesh(replica_devices, mesh.axis_names) - single_replica_sharding = jax.sharding.NamedSharding(replica_mesh, pspec) - - return ocp.type_handlers.SingleReplicaArrayRestoreArgs( - sharding=jax.sharding.NamedSharding(mesh, pspec), - single_replica_sharding=single_replica_sharding, - global_shape=data.shape, - dtype=data.dtype, - ) - - if enable_single_replica_ckpt_restoring: - array_handler = ocp.type_handlers.SingleReplicaArrayHandler( - replica_axis_index=0, - broadcast_memory_limit_bytes=1024 * 1024 * 1000, # 1000 MB limit - ) - ocp.type_handlers.register_type_handler(jax.Array, array_handler, override=True) - is_diloco = bool(maxtext_config and getattr(maxtext_config, "enable_diloco", False)) # Map the expected training state to the on-disk checkpoint dictionary layout: @@ -565,81 +562,45 @@ def map_to_pspec(data): if maxtext_config and getattr(getattr(maxtext_config, "lora", None), "enable_lora", False): restore_target = _filter_lora_trainable_state(restore_target) - restore_args = jax.tree_util.tree_map(map_to_pspec, restore_target) - checkpoint_args = ocp.args.PyTreeRestore( - item=restore_target, - restore_args=restore_args, - partial_restore=True, - ) - match (checkpoint_manager, dataset_type, data_iterator): - # Case 1: Matches if 'checkpoint_manager' is an instance of either EmergencyCheckpointManager - # or EmergencyReplicatorCheckpointManager. The '_' indicates that 'dataset_type' and - # 'data_iterator' can be any value and aren't used in this pattern. - case (checkpoint_manager, _, _) if isinstance( - checkpoint_manager, - ( - EmergencyCheckpointManager, - EmergencyReplicatorCheckpointManager, - ), - ): - restored = checkpoint_manager.restore(step, args=Composite(state=checkpoint_args)).state - if is_diloco: - restored = diloco_checkpoint_utils.from_diloco_checkpoint_dict( - restored, abstract_unboxed_pre_state, config=maxtext_config - ) - elif is_nnx: - restored = _restored_linen_to_nnx(restored, abstract_unboxed_pre_state, config=maxtext_config) - return ( - restored, - None, - ) - # Case 2: Matches if dataset type is "grain" and the data iterator is not a - # PlaceHolderDataIterator and a specific checkpoint file exists for the iterator - case ( - checkpoint_manager, - dataset_type, - data_iterator, - ) if ( - dataset_type == "grain" - and data_iterator - and not isinstance(data_iterator, PlaceHolderDataIterator) - and (checkpoint_manager.directory / str(step) / "iter").exists() - ): - restored, iterator = grain_utility.restore_grain_iterator( - checkpoint_manager, - step, - data_iterator, - checkpoint_args, - expansion_factor_real_data, + # Case 1: emergency / replicator managers restore via their own v0 path. + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + restored = emergency_checkpointing.restore(checkpoint_manager, step, restore_target) + if is_diloco: + restored = diloco_checkpoint_utils.from_diloco_checkpoint_dict( + restored, abstract_unboxed_pre_state, config=maxtext_config ) - if is_diloco: - restored_items = diloco_checkpoint_utils.from_diloco_checkpoint_dict( - restored["items"], abstract_unboxed_pre_state, config=maxtext_config - ) - restored = {"items": restored_items} - elif is_nnx: - restored_items = _restored_linen_to_nnx(restored["items"], abstract_unboxed_pre_state, config=maxtext_config) - restored = {"items": restored_items} - return (restored, iterator) - # 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)) - if is_diloco: - restored_items = diloco_checkpoint_utils.from_diloco_checkpoint_dict( - restored["items"], abstract_unboxed_pre_state, config=maxtext_config - ) - restored = {"items": restored_items} - elif is_nnx: - restored_items = _restored_linen_to_nnx(restored["items"], abstract_unboxed_pre_state, config=maxtext_config) - restored = {"items": restored_items} - return (restored, None) + elif is_nnx: + restored = _restored_linen_to_nnx(restored, abstract_unboxed_pre_state, config=maxtext_config) + return (restored, None) + + # Case 2: standard v1 Checkpointer, restoring the grain iterator in place when + # a "grain" dataset iterator was checkpointed alongside the state. + assert isinstance(checkpoint_manager, ocp.training.Checkpointer) + abstract_checkpointables = {"items": restore_target} + if ( + dataset_type == "grain" + and data_iterator + and not isinstance(data_iterator, PlaceHolderDataIterator) + and (checkpoint_manager.directory / str(step) / "iter").exists() + ): + abstract_checkpointables["iter"] = grain_utility.for_restore( + checkpoint_manager, step, data_iterator, expansion_factor_real_data + ) + restored = checkpoint_manager.load_checkpointables(step, abstract_checkpointables) + if is_diloco: + restored_items = diloco_checkpoint_utils.from_diloco_checkpoint_dict( + restored["items"], abstract_unboxed_pre_state, config=maxtext_config + ) + restored = {"items": restored_items} + elif is_nnx: + restored_items = _restored_linen_to_nnx(restored["items"], abstract_unboxed_pre_state, config=maxtext_config) + restored = {"items": restored_items} + return (restored, None) if source_checkpoint_layout == "safetensors_dynamic": path = load_parameters_from_path or load_full_state_from_path max_logging.log(f"Dynamic On-the-Fly Formatting: Loading SafeTensors from {path}") - # Weights-only for both paths, so the loader gets the weights rather than the whole state: # the HF param mappings name weights, and an NNX state hides them under `model`. params = _abstract_params(abstract_unboxed_pre_state) @@ -655,12 +616,14 @@ def map_to_pspec(data): return restored, restored_params elif load_parameters_from_path != "": params = _abstract_params(abstract_unboxed_pre_state) + restored_params = load_params_from_path( load_parameters_from_path, params, checkpoint_storage_concurrent_gb, use_ocdbt=use_ocdbt, use_zarr3=use_zarr3, + enable_single_replica_ckpt_restoring=bool(enable_single_replica_ckpt_restoring), ) return None, restored_params elif load_full_state_from_path != "": @@ -668,12 +631,12 @@ def map_to_pspec(data): restored_state = _load_full_state_from_path( path=load_full_state_from_path, abstract_unboxed_pre_state=abstract_unboxed_pre_state, - enable_orbax_v1=enable_orbax_v1, checkpoint_conversion_fn=checkpoint_conversion_fn, source_checkpoint_layout=source_checkpoint_layout, checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, use_ocdbt=use_ocdbt, use_zarr3=use_zarr3, + enable_single_replica_ckpt_restoring=bool(enable_single_replica_ckpt_restoring), maxtext_config=maxtext_config, ) return {"items": restored_state}, None @@ -683,22 +646,14 @@ def map_to_pspec(data): def setup_checkpoint_logger(config) -> Any | None: # pytype: disable=attribute-error - """Setup checkpoint logger. - Args: - config - Returns: - CloudLogger - """ - orbax_cloud_logger = None - max_logging.log("Setting up checkpoint logger...") + """DEPRECATED: Setup checkpoint logger.""" + # TODO: b/529622681 - Remove this config option entirely. if config.enable_checkpoint_cloud_logger: - logger_name = f"goodput_{config.run_name}" - orbax_cloud_logger = ocp.logging.CloudLogger( - options=ocp.logging.CloudLoggerOptions(job_name=config.run_name, logger_name=logger_name) + max_logging.warning( + "Cloud logging (enable_checkpoint_cloud_logger) is disabled because" + " Orbax v1 now configures its own logger internally. This config" + " setting is ignored and will be removed." ) - max_logging.log("Successfully set up checkpoint cloud logger.") - - return orbax_cloud_logger def load_params_from_path( @@ -707,11 +662,16 @@ def load_params_from_path( checkpoint_storage_concurrent_gb, use_ocdbt=True, use_zarr3=True, + enable_single_replica_ckpt_restoring: bool = False, ): """Load decode params from checkpoint at specified path.""" assert load_parameters_from_path, "load_parameters_from_path is not defined." max_logging.log(f"restoring params from {load_parameters_from_path}") + # Orbax v1 refuses to read an item subdirectory directly; normalize the documented + # "...//items" form to its checkpoint root and load it by name below. + path = epath.Path(_normalize_checkpoint_root(load_parameters_from_path)) + # On disk the weights live at `params/params/...`: an outer key naming the item, and Flax's # `params` collection inside it. A Linen TrainState.params is that collection; an NNX params # state sits one level below it (bare weights), so wrap it going in and unwrap it coming out. @@ -719,7 +679,7 @@ def load_params_from_path( want = abstract_unboxed_params.to_pure_dict() if is_nnx else abstract_unboxed_params # Determine the restore key based on the leaf directory name to support native and custom SFT - restore_key = os.path.basename(load_parameters_from_path) + restore_key = os.path.basename(str(load_parameters_from_path).rstrip("/")) if restore_key not in ("model_params", "model"): restore_key = "params" @@ -728,39 +688,49 @@ def load_params_from_path( else: params_collection = {"params": want} if is_nnx else want - # *_concurrent_gb should be set for large models, the default is 96. - max_logging.log(f"Creating checkpoint manager with ocdbt={use_ocdbt} and zarr3={use_zarr3}") - 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, - ) - ) - - # This is a memory optimization. We don't want to restore the entire checkpoint - only the params. - # Rather than pass the entire abstract state, which could unnecessarily restore opt_state and such and waste - # memory, we instead specify here that we are just restoring the params field of the checkpoint - # (which itself may be a dictionary containing a key named 'params' or 'model'). - restore_args = ocp.checkpoint_utils.construct_restore_args(params_collection) - restored = ckptr.restore( - epath.Path(load_parameters_from_path), - item={restore_key: params_collection}, - transforms={}, - restore_args={restore_key: restore_args}, + # Memory optimization: restore only the "params" key (the checkpoint may also hold opt_state/step); + # partial_load drops the rest. The abstract carries shape/dtype/sharding directly. + context = checkpoint_context.build_context( + use_ocdbt=use_ocdbt, + use_zarr3=use_zarr3, + checkpoint_storage_concurrent_gb=checkpoint_storage_concurrent_gb, + partial_load=True, + enable_single_replica_ckpt_restoring=enable_single_replica_ckpt_restoring, ) - restored_collection = restored[restore_key] + # Dispatch on the on-disk layout instead of assuming a step root: callers pass step roots, + # v0-style pytree dirs (normalized above), and v0 flat params-only checkpoints + # (save_params_to_path wrote the pytree directly at the directory). + with context: + checkpointable_name = "items" if (path / "items").exists() else None + # Orbax v1 fails a mid-load shape mismatch itself, with an error that reports the + # shapes but not which weight; compare the stored metadata first so the error names + # it. A metadata read failure falls through to the load (worst case: Orbax's error). + try: + stored = ocp.metadata(path, checkpointable_name=checkpointable_name).metadata + except Exception as e: # pylint: disable=broad-except + max_logging.log(f"Skipping pre-load shape check, checkpoint metadata unreadable: {e}") + stored = None + if isinstance(stored, dict): + stored_collection = stored.get(restore_key) + if restore_key == "params" and is_nnx and isinstance(stored_collection, dict): + stored_collection = stored_collection.get("params") + _raise_weight_problems(_weight_mismatches(want, stored_collection, check_missing=False)) + restored = ocp.load( + path, + {restore_key: params_collection}, + checkpointable_name=checkpointable_name, # pyrefly: ignore[bad-argument-type] + ) + restored_collection = restored[restore_key] # pyrefly: ignore[bad-index] + # partial_load lets Orbax return an unmaterialized leaf for a weight the checkpoint lacks, + # and a stored array at its own shape rather than the target's. Either reaches the model and + # fails much later without naming the weight, so check here -- the params-only load + # (load_parameters_path, e.g. SFT) has no init state to fall back on. if restore_key in ("model_params", "model"): restored_weights = restored_collection else: restored_weights = restored_collection["params"] if is_nnx else restored_collection - # `transforms={}` lets Orbax return an unmaterialized leaf for a weight the checkpoint lacks, - # and a stored array at its own shape rather than the target's. Either reaches the model and - # fails much later without naming the weight, so check here -- the params-only load - # (load_parameters_path, e.g. SFT) has no init state to fall back on. _raise_on_weight_mismatch(want, restored_weights) if is_nnx: nnx.replace_by_pure_dict(abstract_unboxed_params, restored_weights) @@ -771,13 +741,19 @@ def load_params_from_path( def save_params_to_path(checkpoint_dir, params, use_ocdbt=True, use_zarr3=True): """Save decode params in checkpoint at specified path.""" assert checkpoint_dir, "checkpoint_dir is not defined." - print(f"Saving quantized params checkpoint with use_ocdbt = {use_ocdbt} and use_zarr3 = {use_zarr3}") - orbax_checkpointer = ocp.PyTreeCheckpointer(use_ocdbt=use_ocdbt, use_zarr3=use_zarr3) - orbax_checkpointer.save(checkpoint_dir, {"params": params}, force=True) - print(f"Quantized params checkpoint saved at: {checkpoint_dir}") + max_logging.log(f"Saving params checkpoint with use_ocdbt={use_ocdbt} and" f" use_zarr3={use_zarr3}") + context = checkpoint_context.build_context(use_ocdbt=use_ocdbt, use_zarr3=use_zarr3) + with context: + ocp.save( + checkpoint_dir, + {"params": params}, + checkpointable_name="items", + overwrite=True, # pyrefly: ignore[bad-argument-type] + ) + max_logging.log(f"Params checkpoint saved at: {checkpoint_dir}") -def load_checkpoint_metadata(checkpoint_dir_path: str) -> dict[str, Any]: +def load_checkpoint_metadata(checkpoint_dir_path: str) -> Any: """Loads custom metadata from an Orbax checkpoint. Args: @@ -787,10 +763,9 @@ def load_checkpoint_metadata(checkpoint_dir_path: str) -> dict[str, Any]: A dictionary containing custom metadata, or an empty dictionary if none is present or loading fails. """ - checkpoint_dir = epath.Path(checkpoint_dir_path) + checkpoint_dir = epath.Path(_normalize_checkpoint_root(checkpoint_dir_path)) try: - ckptr = ocp.StandardCheckpointer() - metadata = ckptr.metadata(checkpoint_dir) + metadata = ocp.checkpointables_metadata(checkpoint_dir) return metadata.custom_metadata or {} except Exception as e: # pylint: disable=broad-except max_logging.log(f"Warning: Failed to load checkpoint metadata: {e}") @@ -888,7 +863,7 @@ def maybe_save_checkpoint(checkpoint_manager, state, config, data_iterator, step # Skip if step directory already exists (e.g. step 0 or prior checkpoints in all_steps()) # to prevent Orbax OCDBT UUID collisions during auto-resume / continuation runs for DiLoCo. - if latest_step(checkpoint_manager) == actual_step or actual_step in checkpoint_manager.all_steps(): + if latest_step(checkpoint_manager) == actual_step or actual_step in all_steps(checkpoint_manager): max_logging.log(f"Checkpoint for step {actual_step} already exists, skipping save.") return @@ -974,62 +949,18 @@ def save_checkpoint(checkpoint_manager, step, state, config=None, data_iterator= f"{step} to finish before starting checkpointing." ) - # specify chunk_byte_size to force orbax to control maximum file size in checkpoint - chunk_byte_size = ( - getattr( - config, - "checkpoint_storage_target_data_file_size_bytes", - DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE, - ) - if config - else DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE - ) - + # LoRA training persists only the adapter weights (plus step/opt_state). if config and getattr(getattr(config, "lora", None), "enable_lora", False): filtered = _filter_lora_trainable_state(state) if filtered: state = filtered - checkpoint_args = ocp.args.PyTreeSave( - item=state, - save_args=jax.tree.map(lambda _: ocp.SaveArgs(chunk_byte_size=chunk_byte_size), state), - ocdbt_target_data_file_size=chunk_byte_size, - ) - save_args_composite = {"items": checkpoint_args} - - if ( - config - and getattr(config, "dataset_type", None) == "grain" - and not isinstance(data_iterator, PlaceHolderDataIterator) - ): - if isinstance(data_iterator, RemoteIteratorWrapper): - # Pass the wrapper directly; GrainCheckpointHandler will call save_state with the step - save_args_composite["iter"] = grain_utility.GrainCheckpointSave( # pyrefly: ignore[bad-assignment] - item=data_iterator - ) # pyrefly: ignore[bad-assignment] - elif not isinstance(data_iterator, list) and isinstance( - data_iterator.local_iterator, ElasticIterator # pyrefly: ignore[missing-attribute] - ): # pyrefly: ignore[missing-attribute] - # ElasticIterator checkpoints a single global scalar shared by all shards. - save_args_composite["iter"] = grain_utility.GrainCheckpointSave( # pyrefly: ignore[bad-assignment] - item=data_iterator.local_iterator - ) # pyrefly: ignore[bad-assignment] - else: - if not isinstance(data_iterator, list): - data_iterator = [data_iterator] - grain_iters_to_save = [] - process_count_total = jax.process_count() * len(data_iterator) - if config.expansion_factor_real_data > 1: - process_count_total = process_count_total // config.expansion_factor_real_data - for i, data_iter in enumerate(data_iterator): - process_index = jax.process_index() + i * jax.process_count() - grain_iters_to_save.append( - (data_iter.local_iterator, process_index, process_count_total) # pyrefly: ignore[missing-attribute] - ) # pyrefly: ignore[missing-attribute] - save_args_composite["iter"] = grain_utility.GrainCheckpointSave( # pyrefly: ignore[bad-assignment] - item=grain_iters_to_save - ) # pyrefly: ignore[bad-assignment] + # Emergency / replicator managers keep the v0 save path. + if isinstance(checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager)): + return emergency_checkpointing.save(checkpoint_manager, step, state, config, force) + # Record config properties needed to validate compatibility at load time + # (e.g. proactive scan_layers verification, LoRA restore). custom_metadata = {} if config: if hasattr(config, "scan_layers"): @@ -1037,13 +968,28 @@ def save_checkpoint(checkpoint_manager, step, state, config=None, data_iterator= if hasattr(config, "lora") and config.lora and getattr(config.lora, "lora_rank", 0) > 0: custom_metadata["lora"] = config.lora.model_dump() - match (checkpoint_manager, config, data_iterator): - case (checkpoint_manager, _, _) if isinstance( - checkpoint_manager, (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager) - ): - emergency_checkpointing.replicator_error_handler(config) - return checkpoint_manager.save(step, args=Composite(state=checkpoint_args), force=force) - case _: - return checkpoint_manager.save( - step, args=Composite(**save_args_composite), force=force, custom_metadata=custom_metadata + # Standard path: Orbax v1 Checkpointer. Storage/chunk options live on the manager's Context. + checkpointables = {"items": state} + if ( + config + and getattr(config, "dataset_type", None) == "grain" + and not isinstance(data_iterator, PlaceHolderDataIterator) + ): + checkpointables["iter"] = grain_utility.for_save(step, data_iterator, config.expansion_factor_real_data) + # The v1 Checkpointer raises for an already-existing step BEFORE consulting the + # save decision policy; v0 silently skipped such saves (should_save ran first). + # Preserve v0 semantics: e.g. resuming from a non-latest step into a directory + # that still holds later/off-interval checkpoints must not kill training. + try: + if getattr(checkpoint_manager, "use_async", False): + # Async save returns once the blocking device-to-host copy is done and writes in + # the background (v0 enable_async_checkpointing parity); a None response means the + # save decision policy declined. Background errors surface on the next save/wait. + response = checkpoint_manager.save_checkpointables_async( + step, checkpointables, force=force, custom_metadata=custom_metadata ) + return response is not None + return checkpoint_manager.save_checkpointables(step, checkpointables, force=force, custom_metadata=custom_metadata) + except FileExistsError as e: # ocp.training StepAlreadyExistsError subclasses FileExistsError + max_logging.log(f"Checkpoint for step {step} already exists, skipping save. ({e})") + return False diff --git a/src/maxtext/common/grain_utility.py b/src/maxtext/common/grain_utility.py index fbd7e8ccc7..630dd20bba 100644 --- a/src/maxtext/common/grain_utility.py +++ b/src/maxtext/common/grain_utility.py @@ -15,17 +15,14 @@ """Grain utility functions for checkpointing.""" import asyncio -import dataclasses import json -from typing import Any, Optional +from typing import Any -from etils import epath import grain from grain import experimental as grain_experimental from grain import python import jax from maxtext.input_pipeline import multihost_dataloading -import numpy as np import orbax.checkpoint as ocp_v0 from orbax.checkpoint import v1 as ocp @@ -60,7 +57,7 @@ async def no_op(): return None -class GrainCheckpointable_v1(ocp.StatefulCheckpointable): +class GrainCheckpointable(ocp.StatefulCheckpointable): """Orbax v1 `StatefulCheckpointable` for MaxText grain data iterators. This is the v1 port of the `GrainCheckpointHandler`: a single object that @@ -89,7 +86,7 @@ class GrainCheckpointable_v1(ocp.StatefulCheckpointable): """ def __init__(self, item, *, restore_process_index=None, restore_process_count=None, step=None): - """Initializes a GrainCheckpointable_v1. + """Initializes a GrainCheckpointable. Args: item: a grain iterator, a grain `ElasticIterator`, a @@ -217,17 +214,17 @@ async def _read_single_fallback(): return _read_single_fallback() -def for_save(step: int, data_iterator: Any, expansion_factor_real_data: int) -> GrainCheckpointable_v1: - """Builds the v1 ``GrainCheckpointable_v1`` for saving the grain iterator.""" +def for_save(step: int, data_iterator: Any, expansion_factor_real_data: int) -> GrainCheckpointable: + """Builds the v1 ``GrainCheckpointable`` for saving the grain iterator.""" if isinstance(data_iterator, RemoteIteratorWrapper): - return GrainCheckpointable_v1(data_iterator, step=step) + return GrainCheckpointable(data_iterator, step=step) if ( not isinstance(data_iterator, list) and hasattr(data_iterator, "local_iterator") and isinstance(data_iterator.local_iterator, ElasticIterator) ): - return GrainCheckpointable_v1(data_iterator.local_iterator) + return GrainCheckpointable(data_iterator.local_iterator) iterators = data_iterator if isinstance(data_iterator, list) else [data_iterator] process_count_total = jax.process_count() * len(iterators) @@ -235,9 +232,7 @@ def for_save(step: int, data_iterator: Any, expansion_factor_real_data: int) -> process_count_total = process_count_total // expansion_factor_real_data if len(iterators) == 1 and process_count_total == jax.process_count(): - return GrainCheckpointable_v1( - iterators[0].local_iterator if hasattr(iterators[0], "local_iterator") else iterators[0] - ) + return GrainCheckpointable(iterators[0].local_iterator if hasattr(iterators[0], "local_iterator") else iterators[0]) specs = [ ( @@ -247,22 +242,22 @@ def for_save(step: int, data_iterator: Any, expansion_factor_real_data: int) -> ) for i, di in enumerate(iterators) ] - return GrainCheckpointable_v1(specs) + return GrainCheckpointable(specs) def for_restore( checkpoint_manager: Any, step: int, data_iterator: Any, expansion_factor_real_data: int -) -> GrainCheckpointable_v1: - """Builds the v1 ``GrainCheckpointable_v1`` for restoring the grain iterator.""" +) -> GrainCheckpointable: + """Builds the v1 ``GrainCheckpointable`` for restoring the grain iterator.""" if isinstance(data_iterator, RemoteIteratorWrapper): - return GrainCheckpointable_v1(data_iterator, step=step) + return GrainCheckpointable(data_iterator, step=step) if ( not isinstance(data_iterator, list) and hasattr(data_iterator, "local_iterator") and isinstance(data_iterator.local_iterator, ElasticIterator) ): - return GrainCheckpointable_v1(data_iterator.local_iterator) + return GrainCheckpointable(data_iterator.local_iterator) directory = checkpoint_manager.directory / str(step) / "iter" process_count_jax = jax.process_count() @@ -281,7 +276,7 @@ def for_restore( ) local_iterators = [x.local_iterator if hasattr(x, "local_iterator") else x for x in data_iterator] restore_process_index = [jax.process_index() + i * process_count_jax for i in range(scaling_factor)] - return GrainCheckpointable_v1( + return GrainCheckpointable( local_iterators, restore_process_index=restore_process_index, restore_process_count=process_count_stored ) @@ -290,7 +285,7 @@ def for_restore( f"{process_count_stored} processes found in Grain checkpoint directory {directory}, matching the number of " "jax processes, please do not set expansion_factor_real_data." ) - return GrainCheckpointable_v1( + return GrainCheckpointable( data_iterator.local_iterator if hasattr(data_iterator, "local_iterator") else data_iterator ) @@ -298,7 +293,7 @@ def for_restore( assert not isinstance( data_iterator, list ), "when expansion_factor_real_data > 1, the data iterator should not be a list." - return GrainCheckpointable_v1( + return GrainCheckpointable( data_iterator.local_iterator if hasattr(data_iterator, "local_iterator") else data_iterator, restore_process_index=jax.process_index(), restore_process_count=process_count_stored, @@ -312,287 +307,3 @@ def for_restore( "https://github.com/AI-Hypercomputer/maxtext/blob/main/docs/guides/data_input_pipeline/" "data_input_grain.md#using-grain" ) - - -# ------------------------------------------------------------------------------ -# TODO(b/532274266): Remove everything below this line once distillation_utils -# supports the new GrainCheckpointHandler. -# ------------------------------------------------------------------------------ - - -class GrainCheckpointHandler(PyGrainCheckpointHandler, ocp_v0.CheckpointHandler): - """A CheckpointHandler that allows specifying process_index and process_count.""" - - def save( - self, - directory: epath.Path, - # `item` is for backwards compatibility with older Orbax API, see - # https://orbax.readthedocs.io/en/latest/guides/checkpoint/api_refactor.html. - item: Optional[Any] = None, - args: Any = None, - ): - """Saves the given iterator to the checkpoint in `directory`.""" - item = item or args.item # pytype:disable=attribute-error - - # RemoteIteratorWrapper handles checkpointing via colocated python - if isinstance(item, RemoteIteratorWrapper): - step = int(directory.parent.name) - item.save_state(step) - return - - # ElasticIterator state is a single global scalar shared by all shards, - # so we write one fixed `process_0.json` from process 0 only. This file - # layout survives changes in `jax.process_count()`. - if isinstance(item, ElasticIterator): - if jax.process_index() == 0: - directory.mkdir(parents=True, exist_ok=True) - filename = directory / "process_0.json" - filename.write_text(json.dumps(item.get_state(), indent=4)) - return - - def save_single_process(item, process_index, process_count): - filename = directory / f"process_{process_index}-of-{process_count}.json" - if isinstance(item, grain.DatasetIterator): - state = json.dumps(item.get_state(), indent=4) - else: - state = item.get_state().decode() - filename.write_text(state) - - if isinstance(item, list): - for local_iterator, process_index, process_count in item: - save_single_process(local_iterator, process_index, process_count) - else: - process_index, process_count = jax.process_index(), jax.process_count() - save_single_process(item, process_index, process_count) - - def restore( - self, - directory: epath.Path, - item: Optional[Any] = None, - args: Any = None, - ) -> Any: - """Restores the given iterator from the checkpoint in `directory`.""" - item = item or args.item - process_index = getattr(args, "process_index", None) - process_count = getattr(args, "process_count", None) - - # In Pathways + colocated_python environment, RemoteIteratorWrapper handles checkpointing - if isinstance(item, RemoteIteratorWrapper): - step = int(directory.parent.name) - item.restore_state(step) - return item - - # McJax and Pathways through controller cases - # ElasticIterator: every process reads the same shared `process_0.json`. - if isinstance(item, ElasticIterator): - filename = directory / "process_0.json" - if not filename.exists(): - raise ValueError(f"File {filename} does not exist.") - item.set_state(json.loads(filename.read_text())) - return item - - def restore_single_process(item, process_index, process_count): - filename = directory / f"process_{process_index}-of-{process_count}.json" - if not filename.exists(): - raise ValueError(f"File {filename} does not exist.") - state = filename.read_text() - if isinstance(item, grain.DatasetIterator): - state = json.loads(state) - else: - state = state.encode() - item.set_state(state) - return item - - if isinstance(item, list): - restored_items = [] - for data_iter, process_idx in zip(item, process_index): # pyrefly: ignore[bad-argument-type] - restored_items.append(restore_single_process(data_iter, process_idx, process_count)) - return restored_items - else: - if process_index is None or process_count is None: - process_index, process_count = jax.process_index(), jax.process_count() - return restore_single_process(item, process_index, process_count) - - -@ocp_v0.args.register_with_handler(GrainCheckpointHandler, for_save=True) -@dataclasses.dataclass -class GrainCheckpointSave(ocp_v0.args.CheckpointArgs): - item: Any - - -@ocp_v0.args.register_with_handler(GrainCheckpointHandler, for_restore=True) -@dataclasses.dataclass -class GrainCheckpointRestore(ocp_v0.args.CheckpointArgs): - item: Any - process_index: Optional[int | list[int]] = None - process_count: Optional[int] = None - - -class GrainCheckpointable(ocp.StatefulCheckpointable): - """Adapts `GrainCheckpointHandler` to Orbax v1's `StatefulCheckpointable`.""" - - def __init__( - self, - *, - save_args: GrainCheckpointSave | None = None, - restore_args: GrainCheckpointRestore | None = None, - ): - self._handler = GrainCheckpointHandler() - self._save_args = save_args - self._restore_args = restore_args - - async def save(self, directory): - """Saves the Grain iterator state to the given directory.""" - # `GrainCheckpointHandler.save` snapshots iterator state (`get_state`) AND - # writes it; both must happen in this (blocking) save phase, NOT in the - # returned background coroutine. - path = await directory.await_creation() - self._handler.save(path, args=self._save_args) - - async def _committed(): # nothing left for the background commit phase - return None - - return _committed() - - async def load(self, directory: epath.Path): - """Loads the Grain iterator state from the given directory.""" - handler, args = self._handler, self._restore_args - - # This will be ran to completion so asynchronous portion is just for - # compatibility with Orbax v1 API. - async def _background_load(): - await asyncio.to_thread(handler.restore, directory, args=args) - - return _background_load() - - -# TODO(b/534897901): Remove find_idx + replica_devices once Orbax exposes multislice. -def find_idx(array: np.ndarray, replica_axis_idx: int): - """Returns the index along given dimension that the current host belongs to.""" - idx = None - for idx, val in np.ndenumerate(array): - if val.process_index == jax.process_index(): - break - return idx[replica_axis_idx] - - -def replica_devices(device_array: np.ndarray, replica_axis_idx: int): - """Returns the devices from the replica that current host belongs to. - - Replicas are assumed to be restricted to the first axis. - - Args: - device_array: devices of the mesh that can be obtained by mesh.devices() - replica_axis_idx: axis dimension along which replica is taken - - Returns: - devices inside the replica that current host is in - """ - idx = find_idx(device_array, replica_axis_idx) - replica_result = np.take(device_array, idx, axis=replica_axis_idx) - return np.expand_dims(replica_result, axis=replica_axis_idx) - - -def prepare_scaled_down_grain_restore_args( - data_iterator: list, process_count_jax: int, process_count_stored: int, directory: epath.Path -) -> GrainCheckpointRestore: - """ - Prepares the restore arguments for a scaled-up (list) data iterator. - - This is used when restoring a checkpoint saved with more processes than - the current run (e.g., 64 files onto 32 JAX processes). - """ - # 1. Validation Assertions - assert isinstance(data_iterator, list), ( - f"{process_count_stored} processes found in Grain checkpoint directory {directory}, but only " - f"{process_count_jax} jax processes in this run, please set expansion_factor_real_data accordingly." - ) - - scaling_factor = len(data_iterator) - expected_process_count = process_count_stored / process_count_jax - assert scaling_factor == expected_process_count, ( - f"Found {process_count_stored} processes in checkpoint and {process_count_jax} " - f"JAX processes, implying a scaling factor of {expected_process_count}. " - f"However, the data_iterator list has {scaling_factor} items." - ) - - # 2. Prepare Arguments - local_iterator_list = [x.local_iterator for x in data_iterator] - # Each JAX process calculates the global indices it's responsible for. - # e.g., process 0 with scaling_factor=2 handles checkpoints from processes [0, 32] - # e.g., process 1 with scaling_factor=2 handles checkpoints from processes [1, 33] - process_index_list = [jax.process_index() + i * process_count_jax for i in range(scaling_factor)] - - return GrainCheckpointRestore(local_iterator_list, process_index=process_index_list, process_count=process_count_stored) - - -def restore_grain_iterator( - checkpoint_manager, - step: int, - data_iterator, - checkpoint_args, - expansion_factor_real_data: int, # This must be defined in the outer scope -) -> tuple[Any, None]: - """ - Handles the complex logic for restoring a Grain data iterator checkpoint. - This function dispatches to the correct restore strategy based on - the number of stored checkpoint files vs. current JAX processes. - """ - if isinstance(data_iterator, RemoteIteratorWrapper): - grain_restore_args = GrainCheckpointRestore(item=data_iterator) - restored_state = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args, iter=grain_restore_args)) - return (restored_state, None) - - # ElasticIterator: one shared `process_0.json` regardless of shard count. - if not isinstance(data_iterator, list) and isinstance(data_iterator.local_iterator, ElasticIterator): - grain_restore_args = GrainCheckpointRestore(item=data_iterator.local_iterator) - restored_state = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args, iter=grain_restore_args)) - return (restored_state, None) - - directory = checkpoint_manager.directory / str(step) / "iter" - process_count_jax = jax.process_count() - - # Count the number of checkpoint files - process_count_stored = len(list(directory.glob("process_*-of-*.json"))) - - grain_restore_args = None - - if process_count_stored > process_count_jax: - # Scaling down from a larger number of hosts. (e.g., 128 files -> 64 processes) - # In this case, each host restores a list of data iterators. - grain_restore_args = prepare_scaled_down_grain_restore_args( - data_iterator, process_count_jax, process_count_stored, directory - ) - - elif process_count_stored == process_count_jax: - # Normal case: number of hosts is the same. (e.g., 64 files -> 64 processes) - assert not isinstance(data_iterator, list), ( - f"{process_count_stored} processes found in Grain checkpoint directory {directory}, matching the number of " - "jax process, please do not set expansion_factor_real_data." - ) - grain_restore_args = GrainCheckpointRestore(data_iterator.local_iterator) - - elif expansion_factor_real_data > 1 and process_count_stored == process_count_jax // expansion_factor_real_data: - # Scaling up to a larger number of hosts.(e.g., 32 files -> 64 processes) - # In this case, a subset of hosts restore the data iterator. - assert not isinstance( - data_iterator, list - ), "when expansion_factor_real_data > 1, the data iterator should not be a list." - grain_restore_args = GrainCheckpointRestore( - data_iterator.local_iterator, process_index=jax.process_index(), process_count=process_count_stored - ) - - else: - # Case 4: Mismatch - raise ValueError( - f"Error restoring Grain checkpoint in {directory}: " - f"The number of stored checkpoint files ({process_count_stored}) " - f"is incompatible with the number of JAX processes ({process_count_jax}). " - "If you are resuming training with a different number of chips, see instructions in " - "https://github.com/AI-Hypercomputer/maxtext/blob/main/docs/guides/data_input_pipeline/" - "data_input_grain.md#using-grain" - ) - - # Call restore once with the composed arguments - restored_state = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args, iter=grain_restore_args)) - return (restored_state, None) diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 269339e056..518ef8fc56 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -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 diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 394e57f514..f13ca2f48b 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -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 @@ -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." @@ -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: diff --git a/src/maxtext/trainers/diloco/utils/spmd_diloco_checkpointing.py b/src/maxtext/trainers/diloco/utils/spmd_diloco_checkpointing.py index 912fdfd6e7..ad60830fb8 100644 --- a/src/maxtext/trainers/diloco/utils/spmd_diloco_checkpointing.py +++ b/src/maxtext/trainers/diloco/utils/spmd_diloco_checkpointing.py @@ -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 @@ -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 "...//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) diff --git a/src/maxtext/trainers/post_train/distillation/distillation_utils.py b/src/maxtext/trainers/post_train/distillation/distillation_utils.py index 3e9df6dbf3..da8b02cfed 100644 --- a/src/maxtext/trainers/post_train/distillation/distillation_utils.py +++ b/src/maxtext/trainers/post_train/distillation/distillation_utils.py @@ -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__( @@ -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) @@ -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 diff --git a/src/maxtext/utils/train_utils.py b/src/maxtext/utils/train_utils.py index 41f74a21b6..a85db28c23 100644 --- a/src/maxtext/utils/train_utils.py +++ b/src/maxtext/utils/train_utils.py @@ -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, @@ -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). diff --git a/tests/integration/diloco_test.py b/tests/integration/diloco_test.py index 3b5650dbf2..71136d367d 100644 --- a/tests/integration/diloco_test.py +++ b/tests/integration/diloco_test.py @@ -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") @@ -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( diff --git a/tests/post_training/unit/lora_utils_test.py b/tests/post_training/unit/lora_utils_test.py index f1c29c4f16..5bbf6f737c 100644 --- a/tests/post_training/unit/lora_utils_test.py +++ b/tests/post_training/unit/lora_utils_test.py @@ -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) @@ -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) @@ -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) @@ -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) @@ -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()) diff --git a/tests/unit/checkpointing_nnx_load_test.py b/tests/unit/checkpointing_nnx_load_test.py index 0c2dae51ff..7ff083821f 100644 --- a/tests/unit/checkpointing_nnx_load_test.py +++ b/tests/unit/checkpointing_nnx_load_test.py @@ -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): @@ -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. @@ -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): @@ -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, @@ -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): @@ -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, @@ -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"]) + 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)) diff --git a/tests/unit/checkpointing_test.py b/tests/unit/checkpointing_test.py index 9440d64821..9abcbcada8 100644 --- a/tests/unit/checkpointing_test.py +++ b/tests/unit/checkpointing_test.py @@ -14,8 +14,6 @@ """Unit tests for the checkpointing components.""" -import asyncio -import json import os from unittest import mock @@ -32,9 +30,9 @@ get_hf_loading_function, ) from maxtext.common import checkpointing -from maxtext.common import grain_utility import numpy as np import optax +import orbax.checkpoint as ocp_v0 import pytest import safetensors.numpy @@ -292,7 +290,6 @@ def __init__(self): load_full_state_from_path="", checkpoint_storage_concurrent_gb=1, abstract_unboxed_pre_state=abstract_state, - enable_orbax_v1=True, source_checkpoint_layout="safetensors_dynamic", maxtext_config=config, ) @@ -308,176 +305,142 @@ def __init__(self): class CheckpointMetadataTest(parameterized.TestCase): """Tests for loading checkpoint custom metadata.""" - @mock.patch.object(checkpointing.ocp, "StandardCheckpointer") - def test_load_checkpoint_metadata(self, mock_checkpointer_cls): - mock_ckptr = mock_checkpointer_cls.return_value + @mock.patch.object(checkpointing.ocp, "checkpointables_metadata") + def test_load_checkpoint_metadata(self, mock_metadata_fn): mock_metadata = mock.MagicMock() mock_metadata.custom_metadata = {"lora": {"lora_rank": 8, "lora_alpha": 16.0}} - mock_ckptr.metadata.return_value = mock_metadata + mock_metadata_fn.return_value = mock_metadata loaded_metadata = checkpointing.load_checkpoint_metadata("dummy/path") self.assertEqual(loaded_metadata.get("lora"), {"lora_rank": 8, "lora_alpha": 16.0}) - mock_ckptr.metadata.assert_called_once() + mock_metadata_fn.assert_called_once() - @mock.patch.object(checkpointing.ocp, "StandardCheckpointer") - def test_load_checkpoint_metadata_handles_exceptions(self, mock_checkpointer_cls): - mock_ckptr = mock_checkpointer_cls.return_value - mock_ckptr.metadata.side_effect = Exception("Checkpoint read error") + @mock.patch.object(checkpointing.ocp, "checkpointables_metadata") + def test_load_checkpoint_metadata_strips_pytree_suffix(self, mock_metadata_fn): + mock_metadata = mock.MagicMock() + mock_metadata.custom_metadata = {"scan_layers": True} + mock_metadata_fn.return_value = mock_metadata + + loaded_metadata = checkpointing.load_checkpoint_metadata("gs://bucket/ckpt/0/items") + self.assertEqual(loaded_metadata, {"scan_layers": True}) + (called_path,) = mock_metadata_fn.call_args.args + self.assertEqual(called_path, epath.Path("gs://bucket/ckpt/0")) + + @mock.patch.object(checkpointing.ocp, "checkpointables_metadata") + def test_load_checkpoint_metadata_handles_exceptions(self, mock_metadata_fn): + mock_metadata_fn.side_effect = Exception("Checkpoint read error") loaded_metadata = checkpointing.load_checkpoint_metadata("corrupt/path") self.assertEqual(loaded_metadata, {}) - mock_ckptr.metadata.assert_called_once() + mock_metadata_fn.assert_called_once() -class GrainCheckpointableEquivalenceTest(parameterized.TestCase): - """Tests to ensure GrainCheckpointable is equivalent to GrainCheckpointHandler.""" +class LoadParamsLayoutCompatTest(parameterized.TestCase): + """load_params_from_path must read every historical params-checkpoint layout.""" def setUp(self): super().setUp() self.tmp_dir = epath.Path(self.create_tempdir().full_path) + self.params = {"dense": {"kernel": jnp.arange(4.0).reshape(2, 2)}} + self.abstract = jax.tree.map(lambda x: jax.ShapeDtypeStruct(x.shape, x.dtype, sharding=x.sharding), self.params) - def test_save_restore_equivalence_single_item(self): - class FakeIterator: - """A fake iterator for testing serialization.""" - - def __init__(self, state=0): - self.state = state + def test_flat_v0_params_checkpoint(self): + """v0 save_params_to_path wrote the pytree FLAT at the directory (no items/ subdir).""" + path = self.tmp_dir / "quantized" + ocp_v0.PyTreeCheckpointer().save(path, {"params": self.params}) - def get_state(self): - return json.dumps({"state": self.state}).encode() + restored = checkpointing.load_params_from_path(str(path), self.abstract, 8) - def set_state(self, state): - self.state = json.loads(state.decode())["state"] + np.testing.assert_allclose(restored["dense"]["kernel"], self.params["dense"]["kernel"]) - def __next__(self): - self.state += 1 - return self.state + def test_step_root_and_items_suffixed_paths(self): + """v1-written step roots load both as the root and as the v0-documented .../items form.""" + root = self.tmp_dir / "0" + checkpointing.save_params_to_path(str(root), self.params) - iterator_v0 = FakeIterator(10) - iterator_v1 = FakeIterator(10) + for path in (str(root), str(root / "items"), str(root / "items") + "/"): + restored = checkpointing.load_params_from_path(path, self.abstract, 8) + np.testing.assert_allclose(restored["dense"]["kernel"], self.params["dense"]["kernel"]) - step = 100 - v0_path = self.tmp_dir / str(step) / "iter_v0" - v1_path = self.tmp_dir / str(step) / "iter_v1" - # v0 Save - handler = grain_utility.GrainCheckpointHandler() - v0_path.mkdir(parents=True, exist_ok=True) - handler.save(v0_path, item=iterator_v0) +class SaveCheckpointStepExistsTest(parameterized.TestCase): + """v0 parity: saving a step that already exists is silently skipped, not fatal.""" - # v1 Save - wrapper = grain_utility.GrainCheckpointable(save_args=grain_utility.GrainCheckpointSave(item=iterator_v1)) + def test_existing_step_returns_false(self): + manager = mock.Mock() + manager.use_async = False + manager.save_checkpointables.side_effect = FileExistsError("step 5 already exists") - class MockDirectory: - """Mock directory for testing checkpointing.""" + saved = checkpointing.save_checkpoint(manager, 5, {"w": 1}) - async def await_creation(self): - v1_path.mkdir(parents=True, exist_ok=True) - return v1_path + self.assertFalse(saved) - commit_func = asyncio.run(wrapper.save(MockDirectory())) - if commit_func: - asyncio.run(commit_func) + def test_existing_step_async_returns_false(self): + manager = mock.Mock() + manager.use_async = True + manager.save_checkpointables_async.side_effect = FileExistsError("step 5 already exists") - # Verify files are identical - v0_file = v0_path / "process_0-of-1.json" - v1_file = v1_path / "process_0-of-1.json" + saved = checkpointing.save_checkpoint(manager, 5, {"w": 1}) - self.assertTrue(v0_file.exists()) - self.assertTrue(v1_file.exists()) - self.assertEqual(v0_file.read_text(), v1_file.read_text()) - - # v0 Restore - restored_iterator_v0 = FakeIterator(0) - args_v0 = grain_utility.GrainCheckpointRestore(item=restored_iterator_v0) - handler.restore(v0_path, args=args_v0) - self.assertEqual(restored_iterator_v0.state, 10) - - # v1 Restore - restored_iterator_v1 = FakeIterator(0) - wrapper_restore = grain_utility.GrainCheckpointable( - restore_args=grain_utility.GrainCheckpointRestore(item=restored_iterator_v1) - ) + self.assertFalse(saved) - load_func = asyncio.run(wrapper_restore.load(v1_path)) - asyncio.run(load_func) - self.assertEqual(restored_iterator_v1.state, 10) - def test_save_restore_equivalence_list_item(self): - class FakeIterator: - """A fake iterator for testing serialization.""" +class SaveCheckpointAsyncTest(parameterized.TestCase): + """save_checkpoint must honor the manager's use_async flag (v0 async parity).""" - def __init__(self, state=0): - self.state = state + def test_async_manager_uses_async_save(self): + manager = mock.Mock() + manager.use_async = True + manager.save_checkpointables_async.return_value = mock.Mock() # AsyncResponse - def get_state(self): - return json.dumps({"state": self.state}).encode() + saved = checkpointing.save_checkpoint(manager, 5, {"w": 1}) - def set_state(self, state): - self.state = json.loads(state.decode())["state"] + self.assertTrue(saved) + manager.save_checkpointables_async.assert_called_once() + manager.save_checkpointables.assert_not_called() - iterator_a = FakeIterator(10) - iterator_b = FakeIterator(20) + def test_async_manager_declined_save_returns_false(self): + manager = mock.Mock() + manager.use_async = True + manager.save_checkpointables_async.return_value = None # decision policy declined - item_v0 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] - item_v1 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] + saved = checkpointing.save_checkpoint(manager, 5, {"w": 1}) - step = 100 - v0_path = self.tmp_dir / str(step) / "iter_v0" - v1_path = self.tmp_dir / str(step) / "iter_v1" + self.assertFalse(saved) - # v0 Save - handler = grain_utility.GrainCheckpointHandler() - v0_path.mkdir(parents=True, exist_ok=True) - handler.save(v0_path, item=item_v0) + def test_sync_manager_uses_blocking_save(self): + manager = mock.Mock() + manager.use_async = False + manager.save_checkpointables.return_value = True - # v1 Save - wrapper = grain_utility.GrainCheckpointable(save_args=grain_utility.GrainCheckpointSave(item=item_v1)) + saved = checkpointing.save_checkpoint(manager, 5, {"w": 1}) - class MockDirectory: - """Mock directory for testing checkpointing.""" + self.assertTrue(saved) + manager.save_checkpointables.assert_called_once() + manager.save_checkpointables_async.assert_not_called() - async def await_creation(self): - v1_path.mkdir(parents=True, exist_ok=True) - return v1_path - commit_func = asyncio.run(wrapper.save(MockDirectory())) - if commit_func: - asyncio.run(commit_func) +class NormalizeCheckpointRootTest(parameterized.TestCase): + """Tests for _normalize_checkpoint_root string manipulation function.""" - # Verify files are identical - v0_file_0 = v0_path / "process_0-of-2.json" - v1_file_0 = v1_path / "process_0-of-2.json" - v0_file_1 = v0_path / "process_1-of-2.json" - v1_file_1 = v1_path / "process_1-of-2.json" + def test_normalize_checkpoint_root(self): + normalize = checkpointing._normalize_checkpoint_root # pylint: disable=protected-access - self.assertTrue(v0_file_0.exists()) - self.assertTrue(v1_file_0.exists()) - self.assertEqual(v0_file_0.read_text(), v1_file_0.read_text()) - - self.assertTrue(v0_file_1.exists()) - self.assertTrue(v1_file_1.exists()) - self.assertEqual(v0_file_1.read_text(), v1_file_1.read_text()) - - # v0 Restore - iterators_restore_v0 = [FakeIterator(0), FakeIterator(0)] - args_v0 = grain_utility.GrainCheckpointRestore(item=iterators_restore_v0, process_index=[0, 1], process_count=2) - handler.restore(v0_path, args=args_v0) - - self.assertEqual(iterators_restore_v0[0].state, 10) - self.assertEqual(iterators_restore_v0[1].state, 20) - - # v1 Restore - iterators_restore_v1 = [FakeIterator(0), FakeIterator(0)] - wrapper_restore = grain_utility.GrainCheckpointable( - restore_args=grain_utility.GrainCheckpointRestore( - item=iterators_restore_v1, process_index=[0, 1], process_count=2 - ) - ) - load_func = asyncio.run(wrapper_restore.load(v1_path)) - asyncio.run(load_func) - self.assertEqual(iterators_restore_v1[0].state, 10) - self.assertEqual(iterators_restore_v1[1].state, 20) + cases = [ + ("/foo/bar/items", "/foo/bar"), + ("/foo/bar/items/", "/foo/bar"), + ("gs://bucket/0/items", "gs://bucket/0"), + ("gs://bucket/0/items/", "gs://bucket/0"), + ("hf://meta-llama/Meta-Llama-3-8B", "hf://meta-llama/Meta-Llama-3-8B"), + ("hf://meta-llama/Meta-Llama-3-8B/", "hf://meta-llama/Meta-Llama-3-8B"), + ("items", "."), + ("items/", "."), + ("an_item", "an_item"), + ("", ""), + ] + for inp, expected in cases: + with self.subTest(inp=inp, expected=expected): + self.assertEqual(normalize(inp), expected) class CheckpointErrorHandlerTest(parameterized.TestCase): diff --git a/tests/unit/elastic_utils_test.py b/tests/unit/elastic_utils_test.py index b795bf9dc4..d6da4b7fb3 100644 --- a/tests/unit/elastic_utils_test.py +++ b/tests/unit/elastic_utils_test.py @@ -563,14 +563,14 @@ def test_checkpoint_exception_guard_checks_scale_up_on_success(self): self.fake_pathwaysutils.is_pathways_backend_used.return_value = True self.fake_manager.available_inactive_slices = {1} # Trigger scale-up - mock_checkpoint_manager = Mock(spec=["wait_until_finished"]) + mock_checkpoint_manager = Mock(spec=checkpointing.ocp.training.Checkpointer) # Successful checkpoint save block raises ScaleUpSignalError to trigger restart with self.assertRaises(ScaleUpSignalError): with checkpointing.checkpoint_exception_guard(config, mock_checkpoint_manager): pass - mock_checkpoint_manager.wait_until_finished.assert_called_once() + mock_checkpoint_manager.wait.assert_called_once() def test_checkpoint_exception_guard_none_manager(self): """Checks that checkpoint_manager=None doesn't raise AttributeError on scale-up.""" diff --git a/tests/unit/grain_utility_test.py b/tests/unit/grain_utility_test.py index 1dd42a4180..7fdc9ad8c5 100644 --- a/tests/unit/grain_utility_test.py +++ b/tests/unit/grain_utility_test.py @@ -15,16 +15,12 @@ """Unit tests for the consolidated grain v1 ``GrainCheckpointable``.""" import asyncio -import json import pathlib import tempfile from typing import Any import unittest from unittest import mock -from absl.testing import absltest -from absl.testing import parameterized -from etils import epath import grain from grain import experimental as grain_experimental import grain.sharding @@ -33,7 +29,7 @@ ElasticIterator = grain_experimental.ElasticIterator -GrainCheckpointable = grain_utility.GrainCheckpointable_v1 +GrainCheckpointable = grain_utility.GrainCheckpointable def _std_iter() -> Any: @@ -188,156 +184,3 @@ def test_for_restore_elastic(self): # pylint: enable=protected-access - - -# ------------------------------------------------------------------------------ -# TODO(b/532274266): Remove everything below this line once distillation_utils -# supports the new GrainCheckpointHandler. -# ------------------------------------------------------------------------------ - - -class GrainCheckpointableEquivalenceTest(parameterized.TestCase): - """Tests to ensure GrainCheckpointable is equivalent to GrainCheckpointHandler.""" - - def setUp(self): - super().setUp() - self.tmp_dir = epath.Path(self.create_tempdir().full_path) - - def test_save_restore_equivalence_single_item(self): - class FakeIterator: - """A fake iterator for testing.""" - - def __init__(self, state=0): - self.state = state - - def get_state(self): - return json.dumps({"state": self.state}).encode() - - def set_state(self, state): - self.state = json.loads(state.decode())["state"] - - def __next__(self): - self.state += 1 - return self.state - - iterator_v0 = FakeIterator(10) - iterator_v1 = FakeIterator(10) - - step = 100 - v0_path = self.tmp_dir / str(step) / "iter_v0" - v1_path = self.tmp_dir / str(step) / "iter_v1" - - # v0 Save - handler = grain_utility.GrainCheckpointHandler() - v0_path.mkdir(parents=True, exist_ok=True) - handler.save(v0_path, item=iterator_v0) - - # v1 Save - wrapper = GrainCheckpointable(iterator_v1) - - class MockDirectory: - - async def await_creation(self): - v1_path.mkdir(parents=True, exist_ok=True) - return v1_path - - commit_func = asyncio.run(wrapper.save(MockDirectory())) - if commit_func: - asyncio.run(commit_func) - - # Verify files are identical - v0_file = v0_path / "process_0-of-1.json" - v1_file = v1_path / "process_0-of-1.json" - - self.assertTrue(v0_file.exists()) - self.assertTrue(v1_file.exists()) - self.assertEqual(v0_file.read_text(), v1_file.read_text()) - - # v0 Restore - restored_iterator_v0 = FakeIterator(0) - args_v0 = grain_utility.GrainCheckpointRestore(item=restored_iterator_v0) - handler.restore(v0_path, args=args_v0) - self.assertEqual(restored_iterator_v0.state, 10) - - # v1 Restore - restored_iterator_v1 = FakeIterator(0) - wrapper_restore = GrainCheckpointable(restored_iterator_v1) - - load_func = asyncio.run(wrapper_restore.load(v1_path)) - asyncio.run(load_func) - self.assertEqual(restored_iterator_v1.state, 10) - - def test_save_restore_equivalence_list_item(self): - class FakeIterator: - """A fake iterator for testing.""" - - def __init__(self, state=0): - self.state = state - - def get_state(self): - return json.dumps({"state": self.state}).encode() - - def set_state(self, state): - self.state = json.loads(state.decode())["state"] - - iterator_a = FakeIterator(10) - iterator_b = FakeIterator(20) - - item_v0 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] - item_v1 = [(iterator_a, 0, 2), (iterator_b, 1, 2)] - - step = 100 - v0_path = self.tmp_dir / str(step) / "iter_v0" - v1_path = self.tmp_dir / str(step) / "iter_v1" - - # v0 Save - handler = grain_utility.GrainCheckpointHandler() - v0_path.mkdir(parents=True, exist_ok=True) - handler.save(v0_path, item=item_v0) - - # v1 Save - wrapper = GrainCheckpointable(item_v1) - - class MockDirectory: - - async def await_creation(self): - v1_path.mkdir(parents=True, exist_ok=True) - return v1_path - - commit_func = asyncio.run(wrapper.save(MockDirectory())) - if commit_func: - asyncio.run(commit_func) - - # Verify files are identical - v0_file_0 = v0_path / "process_0-of-2.json" - v1_file_0 = v1_path / "process_0-of-2.json" - v0_file_1 = v0_path / "process_1-of-2.json" - v1_file_1 = v1_path / "process_1-of-2.json" - - self.assertTrue(v0_file_0.exists()) - self.assertTrue(v1_file_0.exists()) - self.assertEqual(v0_file_0.read_text(), v1_file_0.read_text()) - - self.assertTrue(v0_file_1.exists()) - self.assertTrue(v1_file_1.exists()) - self.assertEqual(v0_file_1.read_text(), v1_file_1.read_text()) - - # v0 Restore - iterators_restore_v0 = [FakeIterator(0), FakeIterator(0)] - args_v0 = grain_utility.GrainCheckpointRestore(item=iterators_restore_v0, process_index=[0, 1], process_count=2) - handler.restore(v0_path, args=args_v0) - - self.assertEqual(iterators_restore_v0[0].state, 10) - self.assertEqual(iterators_restore_v0[1].state, 20) - - # v1 Restore - iterators_restore_v1 = [FakeIterator(0), FakeIterator(0)] - wrapper_restore = GrainCheckpointable(iterators_restore_v1, restore_process_index=[0, 1], restore_process_count=2) - load_func = asyncio.run(wrapper_restore.load(v1_path)) - asyncio.run(load_func) - self.assertEqual(iterators_restore_v1[0].state, 10) - self.assertEqual(iterators_restore_v1[1].state, 20) - - -if __name__ == "__main__": - absltest.main() diff --git a/tests/unit/train_state_nnx_checkpoint_test.py b/tests/unit/train_state_nnx_checkpoint_test.py index 172bd09a72..ba44a46438 100644 --- a/tests/unit/train_state_nnx_checkpoint_test.py +++ b/tests/unit/train_state_nnx_checkpoint_test.py @@ -396,7 +396,7 @@ def _invoke_maybe_save(self, state, pure_nnx): # checkpoint_period=1 keeps force_ckpt_save False regardless of actual_step. config = self._config(pure_nnx=pure_nnx, checkpoint_period=1) mgr = mock.MagicMock() - mgr.reached_preemption.return_value = False + mgr.latest = None # no existing checkpoint -> _latest_step() is None, so the save proceeds captured = {} @@ -405,7 +405,10 @@ def fake_save_checkpoint(_mgr, step, state_arg, *_args, **_kwargs): captured["state"] = state_arg return False # no save happened => print_save_message is skipped - with mock.patch.object(checkpointing, "save_checkpoint", side_effect=fake_save_checkpoint): + with ( + mock.patch.object(checkpointing, "save_checkpoint", side_effect=fake_save_checkpoint), + mock.patch.object(checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False), + ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=None) return captured @@ -437,34 +440,24 @@ def test_nnx_state_is_saved_in_linen_layout(self): config = self._config(pure_nnx=True, enable_checkpointing=True, checkpoint_period=1) mgr = mock.MagicMock() - mgr.reached_preemption.return_value = False - - captured = {} - def fake_save(_step, *args, **kwargs): - composite = kwargs.get("args") - if composite: - for key in ["items", "state"]: - if hasattr(composite, "_items") and key in composite._items: # pylint: disable=protected-access - val = composite[key] - if val is not None and hasattr(val, "item"): - captured["state"] = val.item - break - return True + mgr.use_async = False - mgr.save.side_effect = fake_save - checkpointing.save_checkpoint(mgr, self.N_STEPS - 1, state, config=config, force=True) + with mock.patch.object(checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False): + checkpointing.save_checkpoint(mgr, self.N_STEPS - 1, state, config=config, force=True) # save_checkpoint should pass a plain dict in Linen layout to Orbax, not the nnx.State. - self.assertIsInstance(captured["state"], dict) - self.assertNotIsInstance(captured["state"], nnx.State) + mgr.save_checkpointables.assert_called_once() + saved = mgr.save_checkpointables.call_args.args[1]["items"] + self.assertIsInstance(saved, dict) + self.assertNotIsInstance(saved, nnx.State) # Linen layout: {params: {params: ...}, step, opt_state}; not the NNX {model, optimizer}. - self.assertIn("params", captured["state"]) - self.assertIn("step", captured["state"]) - self.assertIn("opt_state", captured["state"]) - self.assertNotIn("model", captured["state"]) - self.assertNotIn("optimizer", captured["state"]) - self.assertIn("params", captured["state"]["params"]) + self.assertIn("params", saved) + self.assertIn("step", saved) + self.assertIn("opt_state", saved) + self.assertNotIn("model", saved) + self.assertNotIn("optimizer", saved) + self.assertIn("params", saved["params"]) def test_linen_state_is_passed_through_unchanged(self): """For pure_nnx=False, maybe_save_checkpoint must pass the original TrainState object through.""" @@ -480,13 +473,15 @@ def test_maybe_save_checkpoint_skips_if_already_saved(self): config = self._config(checkpoint_period=1) mgr = mock.MagicMock() - mgr.reached_preemption.return_value = False - # Mock latest_step to return the same actual_step - mgr.latest_step.return_value = actual_step + # Latest saved step matches actual_step -> save should be skipped. + mgr.latest = mock.MagicMock(step=actual_step) save_checkpoint_mock = mock.MagicMock() - with mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock): + with ( + mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock), + mock.patch.object(checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False), + ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=None) # Assert that save_checkpoint was NOT called! @@ -499,14 +494,16 @@ def test_maybe_save_checkpoint_saves_if_not_already_saved(self): config = self._config(checkpoint_period=1) mgr = mock.MagicMock() - mgr.reached_preemption.return_value = False - # Mock latest_step to return a different step (or None) - mgr.latest_step.return_value = actual_step - 1 + # Latest saved step differs from actual_step -> save should happen. + mgr.latest = mock.MagicMock(step=actual_step - 1) save_checkpoint_mock = mock.MagicMock() save_checkpoint_mock.return_value = False - with mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock): + with ( + mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock), + mock.patch.object(checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False), + ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=None) # Assert that save_checkpoint WAS called! @@ -519,17 +516,18 @@ def test_maybe_save_checkpoint_skips_non_checkpoint_step_before_state_work( state = mock.Mock() config = self._config() mgr = mock.MagicMock() - mgr.reached_preemption.return_value = False with ( mock.patch.object(checkpointing, "save_checkpoint") as save_checkpoint_mock, mock.patch.object(train_state_nnx, "to_checkpoint_dict") as to_checkpoint_dict_mock, + mock.patch.object( + checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False + ) as reached_preemption_mock, ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=3) - mgr.latest_step.assert_not_called() - mgr.reached_preemption.assert_called_once_with(3) - mgr.wait_until_finished.assert_not_called() + reached_preemption_mock.assert_called_once_with(3) + mgr.wait.assert_not_called() to_checkpoint_dict_mock.assert_not_called() save_checkpoint_mock.assert_not_called() @@ -540,18 +538,19 @@ def test_maybe_save_checkpoint_handles_preemption_on_non_checkpoint_step( state = mock.Mock() config = self._config() mgr = mock.MagicMock() - mgr.reached_preemption.return_value = True with ( mock.patch.object(checkpointing, "save_checkpoint") as save_checkpoint_mock, mock.patch.object(train_state_nnx, "to_checkpoint_dict") as to_checkpoint_dict_mock, + mock.patch.object( + checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=True + ) as reached_preemption_mock, ): with self.assertRaises(checkpointing.exceptions.StopTraining): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=3) - mgr.latest_step.assert_not_called() - mgr.reached_preemption.assert_called_once_with(3) - mgr.wait_until_finished.assert_called_once_with() + reached_preemption_mock.assert_called_once_with(3) + mgr.wait.assert_called_once_with() to_checkpoint_dict_mock.assert_not_called() save_checkpoint_mock.assert_not_called() @@ -569,16 +568,19 @@ def test_maybe_save_checkpoint_allows_local_checkpoint_period(self): **{checkpoint_flag: True}, ) mgr = mock.MagicMock() - mgr.latest_step.return_value = None - mgr.reached_preemption.return_value = False + mgr.latest = None save_checkpoint_mock = mock.MagicMock(return_value=False) - with mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock): + with ( + mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock), + mock.patch.object( + checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False + ) as reached_preemption_mock, + ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=5) - mgr.latest_step.assert_called_once_with() - mgr.reached_preemption.assert_called_once_with(5) - mgr.wait_until_finished.assert_not_called() + reached_preemption_mock.assert_called_once_with(5) + mgr.wait.assert_not_called() save_checkpoint_mock.assert_called_once() def test_maybe_save_checkpoint_allows_mtc_period_with_continuous_policy( @@ -594,16 +596,19 @@ def test_maybe_save_checkpoint_allows_mtc_period_with_continuous_policy( ) mgr = mock.MagicMock() mgr.should_save.return_value = False - mgr.latest_step.return_value = None - mgr.reached_preemption.return_value = False + mgr.latest = None save_checkpoint_mock = mock.MagicMock(return_value=False) - with mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock): + with ( + mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock), + mock.patch.object( + checkpointing.multihost_utils, "reached_preemption_sync_point", return_value=False + ) as reached_preemption_mock, + ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=5) mgr.should_save.assert_called_once_with(5) - mgr.latest_step.assert_called_once_with() - mgr.reached_preemption.assert_called_once_with(5) + reached_preemption_mock.assert_called_once_with(5) save_checkpoint_mock.assert_called_once() def test_maybe_save_checkpoint_checks_scale_up_after_unsaved_dispatch(self): @@ -611,16 +616,21 @@ def test_maybe_save_checkpoint_checks_scale_up_after_unsaved_dispatch(self): state = mock.Mock() config = self._config(checkpoint_period=1, elastic_enabled=True) mgr = mock.MagicMock() - mgr.latest_step.return_value = None - mgr.reached_preemption.return_value = False + mgr.latest = None save_checkpoint_mock = mock.MagicMock(return_value=False) with ( mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock), mock.patch.object(checkpointing.elastic_utils, "maybe_elastic_scale_up") as mock_maybe_scale_up, + mock.patch.object( + checkpointing.multihost_utils, + "reached_preemption_sync_point", + return_value=False, + ) as reached_preemption_mock, ): checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=5) + reached_preemption_mock.assert_called_once_with(5) save_checkpoint_mock.assert_called_once() mock_maybe_scale_up.assert_called_once_with(config, mgr) diff --git a/tests/unit/train_utils_test.py b/tests/unit/train_utils_test.py index 30f6159257..9dd4f36084 100644 --- a/tests/unit/train_utils_test.py +++ b/tests/unit/train_utils_test.py @@ -131,12 +131,12 @@ def test_grain_checkpoint_round_trip_through_reorder_view(self): view = train_utils._ReorderedDataIterator(lambda batch: batch * 10, iterator) self.assertEqual(next(view), 0) self.assertEqual(next(view), 10) - ocp.save_checkpointables(str(path), {"iter": grain_utility.GrainCheckpointable_v1(iterator)}) + ocp.save_checkpointables(str(path), {"iter": grain_utility.GrainCheckpointable(iterator)}) self.assertTrue((path / "iter" / "process_0-of-1.json").exists()) restored = iter(grain.MapDataset.range(10).to_iter_dataset()) restored_view = train_utils._ReorderedDataIterator(lambda batch: batch * 10, restored) - ocp.load_checkpointables(str(path), {"iter": grain_utility.GrainCheckpointable_v1(restored)}) + ocp.load_checkpointables(str(path), {"iter": grain_utility.GrainCheckpointable(restored)}) self.assertEqual(next(restored_view), 20) def test_eval_view_consume_reset_consume(self):