From b7b4dd41df0992b159b20e2e9c1a8787f6bb6d66 Mon Sep 17 00:00:00 2001 From: Sadi Kneipp Date: Fri, 14 Aug 2026 18:01:39 +0000 Subject: [PATCH] Fix snapshot restore for abstract shapes and memory kind (b/545734455) --- .../requirements/requirements.txt | 2 +- src/maxtext/trainers/pre_train/train.py | 31 +++++++++++++++---- src/maxtext/utils/elastic_utils.py | 2 +- src/maxtext/utils/train_utils.py | 16 +++++++--- 4 files changed, 39 insertions(+), 12 deletions(-) diff --git a/src/dependencies/requirements/requirements.txt b/src/dependencies/requirements/requirements.txt index 4bef34952e..296c03c55e 100644 --- a/src/dependencies/requirements/requirements.txt +++ b/src/dependencies/requirements/requirements.txt @@ -23,7 +23,7 @@ numpy omegaconf optax orbax-checkpoint -git+https://github.com/AI-Hypercomputer/pathways-utils.git@ce843bad8198fe958c22c22feee6b47e46dbf036 +git+https://github.com/AI-Hypercomputer/pathways-utils.git@ecd319c096ff8b926df35d82d6603547ec2aa0cf pillow pre-commit protobuf diff --git a/src/maxtext/trainers/pre_train/train.py b/src/maxtext/trainers/pre_train/train.py index 4e243ea870..cea983e6f8 100644 --- a/src/maxtext/trainers/pre_train/train.py +++ b/src/maxtext/trainers/pre_train/train.py @@ -24,6 +24,7 @@ import os import sys import logging +import time from absl import app import optax @@ -970,8 +971,12 @@ def recover( "optimizer": nnx.to_pure_dict(nnx.state(active_state.optimizer)), } restored_dict = jax.device_put(active_dict, sharding_dict) - nnx.update(state.model, restored_dict["model"]) - nnx.update(state.optimizer, restored_dict["optimizer"]) + m_state = nnx.state(state.model) + nnx.replace_by_pure_dict(m_state, restored_dict["model"]) + nnx.update(state.model, m_state) + opt_state = nnx.state(state.optimizer) + nnx.replace_by_pure_dict(opt_state, restored_dict["optimizer"]) + nnx.update(state.optimizer, opt_state) restored_state = state restored_step = int(state.optimizer.step.value) _logger.info( @@ -1006,8 +1011,18 @@ def recover( replicated_abstract_dict = train_utils.replicate_single_device_sharded_arrays(abstract_dict) restored_dict = snapshot_mgr.load(replicated_abstract_dict) restored_dict = train_utils.restore_original_shardings(restored_dict, abstract_dict) - nnx.update(state.model, restored_dict["model"]) - nnx.update(state.optimizer, restored_dict["optimizer"]) + merged = jax.tree.map( + lambda ckpt, init: init if isinstance(ckpt, jax.ShapeDtypeStruct) else ckpt, + restored_dict, + abstract_dict, + is_leaf=lambda x: isinstance(x, jax.ShapeDtypeStruct), + ) + m_state = nnx.state(state.model) + nnx.replace_by_pure_dict(m_state, merged["model"]) + nnx.update(state.model, m_state) + opt_state = nnx.state(state.optimizer) + nnx.replace_by_pure_dict(opt_state, merged["optimizer"]) + nnx.update(state.optimizer, opt_state) restored_state = state if metric_logger_instance is not None: @@ -1043,10 +1058,14 @@ def recover( ) break - except pathways_manager.ScaleUpSignalError as e: + except (pathways_manager.ScaleUpSignalError, jax.errors.JaxRuntimeError) as e: + if isinstance(e, jax.errors.JaxRuntimeError) and not elastic.is_error_due_to_slice_down(e): + raise _logger.info( - "ScaleUpSignalError caught during recovery: %s. Retrying recovery.", e + "Transient slice failure / scale event caught during recovery: %s. Retrying recovery.", e ) + active_state = None + time.sleep(2) def train_loop(config, recorder, state=None): diff --git a/src/maxtext/utils/elastic_utils.py b/src/maxtext/utils/elastic_utils.py index 3d6e1e321d..13ed182f7e 100644 --- a/src/maxtext/utils/elastic_utils.py +++ b/src/maxtext/utils/elastic_utils.py @@ -173,7 +173,7 @@ def ensure_elastic_manager_initialized(config): timeout = config.elastic_timeout_seconds max_logging.log(f"[*] Waiting for {min_slices} slices to be active before initializing config...") all_active_slices = elastic.wait_for_slices( - slice_count=len(slice_to_devices), # Temporary for tests, min_slices, + slice_count=min_slices, slice_to_devices=slice_to_devices, timeout=timeout, ) diff --git a/src/maxtext/utils/train_utils.py b/src/maxtext/utils/train_utils.py index 7913edbe2b..11b9a0e52d 100644 --- a/src/maxtext/utils/train_utils.py +++ b/src/maxtext/utils/train_utils.py @@ -517,10 +517,18 @@ def _replicate(x): def restore_original_shardings(restored_pytree, original_abstract_pytree): """Puts restored state back onto the original abstract state shardings.""" def _put(restored_leaf, abstract_leaf): - if isinstance(restored_leaf, jax.Array) and isinstance( - abstract_leaf, (jax.Array, jax.ShapeDtypeStruct) - ): - if restored_leaf.sharding != abstract_leaf.sharding: + if hasattr(restored_leaf, "sharding") and hasattr(abstract_leaf, "sharding"): + if restored_leaf.sharding != abstract_leaf.sharding or ( + hasattr(restored_leaf.sharding, "memory_kind") + and hasattr(abstract_leaf.sharding, "memory_kind") + and restored_leaf.sharding.memory_kind != abstract_leaf.sharding.memory_kind + ): + if isinstance(restored_leaf, jax.ShapeDtypeStruct): + return jax.ShapeDtypeStruct( + restored_leaf.shape, + restored_leaf.dtype, + sharding=abstract_leaf.sharding, + ) return jax.device_put(restored_leaf, abstract_leaf.sharding) return restored_leaf