1515
1616"""Create an Orbax CheckpointManager with specified (Async or not) Checkpointer."""
1717
18+ import datetime
1819import time
1920from typing import Any
2021
21- from absl import flags
22- import datetime
2322from etils import epath
2423from flax import nnx
2524from flax .training import train_state
25+ from grain .experimental import ElasticIterator
2626import jax
27- from maxtext .utils .globals import DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE
27+ from maxtext .checkpoint_conversion .utils .load_dynamic import load_safetensors_dynamic_state
28+ from maxtext .common import emergency_checkpointing
29+ from maxtext .common import grain_utility
30+ from maxtext .common import train_state_nnx
2831from maxtext .input_pipeline .multihost_dataloading import MultiHostDataLoadIterator
2932from maxtext .input_pipeline .multihost_dataloading import RemoteIteratorWrapper
3033from maxtext .input_pipeline .synthetic_data_processing import PlaceHolderDataIterator
31- from maxtext .common import grain_utility
32- from maxtext .common import train_state_nnx
34+ from maxtext .utils import elastic_utils
3335from maxtext .utils import exceptions
34- from maxtext .utils import max_logging
3536from maxtext .utils import gcs_utils
36- from maxtext .utils import elastic_utils
37- from maxtext .checkpoint_conversion .utils .load_dynamic import load_safetensors_dynamic_state
38-
37+ from maxtext .utils import max_logging
38+ from maxtext .utils .globals import DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE
3939import orbax .checkpoint as ocp
4040from orbax .checkpoint import v1 as ocp_v1
4141from orbax .checkpoint ._src .arrays import sharding as sharding_utils
4242from orbax .checkpoint ._src .checkpoint_managers import preservation_policy as preservation_policy_lib
4343from orbax .checkpoint ._src .checkpoint_managers import save_decision_policy as save_decision_policy_lib
44- import orbax .checkpoint .experimental .emergency .checkpoint_manager as emergency_checkpoint_manager
45- import orbax .checkpoint .experimental .emergency .replicator_checkpoint_manager as emergency_replicator_checkpoint_manager
46- # pylint: disable=too-many-positional-arguments
4744
48- from grain .experimental import ElasticIterator
4945
50- CheckpointManager = ocp .CheckpointManager
5146CheckpointManagerOptions = ocp .CheckpointManagerOptions
5247Composite = ocp .args .Composite
5348PyTreeCheckpointHandler = ocp .PyTreeCheckpointHandler
54- EmergencyCheckpointManager = emergency_checkpoint_manager .CheckpointManager
55- LocalCheckpointOptions = emergency_checkpoint_manager .LocalCheckpointOptions
56- PersistentCheckpointOptions = emergency_checkpoint_manager .PersistentCheckpointOptions
57- EmergencyReplicatorCheckpointManager = emergency_replicator_checkpoint_manager .ReplicatorCheckpointManager
49+ # Backward compatibility aliases for v0 emergency managers.
50+ EmergencyCheckpointManager = emergency_checkpointing .CheckpointManager
51+ EmergencyReplicatorCheckpointManager = emergency_checkpointing .ReplicatorCheckpointManager
52+ create_orbax_emergency_checkpoint_manager = emergency_checkpointing .create_emergency_checkpoint_manager
53+ create_orbax_emergency_replicator_checkpoint_manager = emergency_checkpointing .create_replicator_checkpoint_manager
54+
55+ # Union of CheckpointManager / the emergency factories return; used in type hints.
56+ CheckpointManager = ocp .CheckpointManager | EmergencyCheckpointManager | EmergencyReplicatorCheckpointManager
5857
5958
6059def _weight_mismatches (want , have , path = ()):
@@ -336,7 +335,7 @@ def create_orbax_checkpoint_manager(
336335 async_options = ocp .AsyncOptions (
337336 timeout_secs = int (datetime .timedelta (minutes = 60 ).total_seconds ()),
338337 )
339- manager = CheckpointManager (
338+ manager = ocp . CheckpointManager (
340339 p ,
341340 item_names = item_names ,
342341 item_handlers = item_handlers ,
@@ -356,115 +355,6 @@ def create_orbax_checkpoint_manager(
356355 return manager
357356
358357
359- def create_orbax_emergency_checkpoint_manager (
360- local_checkpoint_dir : str ,
361- persistent_checkpoint_dir : str ,
362- global_mesh : jax .sharding .Mesh ,
363- abstract_state : Any ,
364- local_save_interval_steps : int ,
365- persistent_save_interval_steps : int ,
366- orbax_logger : Any = None , # pytype: disable=attribute-error
367- ):
368- """Returns an emergency checkpoint manager."""
369- flags .FLAGS .experimental_orbax_use_distributed_process_id = True
370- max_logging .log ("Creating emergency checkpoint manager..." )
371-
372- # Only create local directories if running on GPUs as the previous directory structure might be assumed by TPUs.
373- if global_mesh .devices .flatten ()[0 ].platform == "gpu" :
374- # pylint: disable=protected-access
375- local_checkpoint_dir = f"{ local_checkpoint_dir } /{ jax ._src .distributed .global_state .process_id } "
376- local_p = epath .Path (local_checkpoint_dir )
377- local_p .mkdir (exist_ok = True , parents = True )
378-
379- persistent_p = gcs_utils .mkdir_and_check_permissions (persistent_checkpoint_dir )
380-
381- # pure_nnx saves via to_checkpoint_dict (Linen params/opt_state/step plus an nnx_aux
382- # subtree), but the emergency manager restores against the abstract it is built with.
383- # Convert it the same way so it matches what is on disk; restore reshapes back to NNX.
384- if isinstance (abstract_state , nnx .State ):
385- abstract_state = train_state_nnx .to_checkpoint_dict (abstract_state )
386-
387- manager = EmergencyCheckpointManager (
388- local_checkpoint_dir ,
389- persistent_p ,
390- global_mesh = global_mesh ,
391- abstract_state = abstract_state ,
392- options = emergency_checkpoint_manager .CheckpointManagerOptions (
393- local = LocalCheckpointOptions (save_interval_steps = local_save_interval_steps ),
394- persistent = PersistentCheckpointOptions (save_interval_steps = persistent_save_interval_steps ),
395- ),
396- logger = orbax_logger ,
397- )
398-
399- max_logging .log ("Emergency checkpoint manager created!" )
400- return manager
401-
402-
403- def create_orbax_emergency_replicator_checkpoint_manager (
404- local_checkpoint_dir : str ,
405- save_interval_steps : int ,
406- global_mesh : jax .sharding .Mesh ,
407- colocated_python_checkpointing : bool = False ,
408- ):
409- """Returns an emergency replicator checkpoint manager."""
410- flags .FLAGS .experimental_orbax_use_distributed_process_id = True
411- max_logging .log ("Creating emergency replicator checkpoint manager..." )
412-
413- manager = EmergencyReplicatorCheckpointManager (
414- epath .Path (local_checkpoint_dir ),
415- options = emergency_replicator_checkpoint_manager .ReplicatorCheckpointManagerOptions (
416- save_interval_steps = save_interval_steps ,
417- use_colocated_python = colocated_python_checkpointing ,
418- ),
419- global_mesh = global_mesh ,
420- )
421-
422- max_logging .log ("Emergency replicator checkpoint manager created!" )
423- return manager
424-
425-
426- def replicator_error_handler (config : Any ):
427- """Replicator error handler to handle errors in replicator service."""
428- if config .enable_multi_tier_checkpointing :
429- local_dir = config .local_checkpoint_directory
430- replicator_errors_file = f"{ local_dir } /replicator.errors"
431- replicator_failed_file = f"{ local_dir } /replicator.failed"
432- process_replicator_error_file (replicator_errors_file )
433-
434- # if the replicator.failed file exists, then we have a fatal error
435- is_fatal = process_replicator_error_file (replicator_failed_file )
436- if is_fatal :
437- raise ValueError ("Replicator fatal error found in replicator.failed file." )
438-
439-
440- def process_replicator_error_file (error_file : str ) -> bool :
441- """Handles replicator errors by reading, logging, cleaning the error file."""
442- error_file_path_exists = epath .Path (error_file ).exists ()
443- if error_file_path_exists :
444- max_logging .log (f"replicator_error_handler: file found: { error_file } ." )
445- read_replicator_error_file (error_file )
446- cleanup_replicator_error_file (error_file )
447-
448- return error_file_path_exists
449-
450-
451- def read_replicator_error_file (error_file : str ):
452- """Read replicator errors file."""
453- try :
454- error_data = epath .Path (error_file ).read_text ()
455- max_logging .log (f"Contents of replicator error file:\n { error_data } " )
456- except (OSError , ValueError ) as e :
457- max_logging .log ("replicator_error_handler: Failed to read contents of failed" f" file: { e } " )
458-
459-
460- def cleanup_replicator_error_file (error_file : str ):
461- """Clean up replicator errors file."""
462- try :
463- epath .Path (error_file ).unlink ()
464- except (OSError , ValueError ) as e :
465- max_logging .log ("replicator_error_handler: Failed to remove replicator errors file:" f" { e } " )
466-
467-
468358def print_save_message (step , async_checkpointing ):
469359 if async_checkpointing :
470360 max_logging .log (f"Started an asynchronous checkpoint save for step { step } " )
@@ -944,7 +834,7 @@ def save_checkpoint(checkpoint_manager, step, state, config=None, data_iterator=
944834 case (checkpoint_manager , _, _) if isinstance (
945835 checkpoint_manager , (EmergencyCheckpointManager , EmergencyReplicatorCheckpointManager )
946836 ):
947- replicator_error_handler (config )
837+ emergency_checkpointing . replicator_error_handler (config )
948838 return checkpoint_manager .save (step , args = Composite (state = checkpoint_args ), force = force )
949839 case _:
950840 return checkpoint_manager .save (
0 commit comments