From 4efbf8d57452865e9803d9bc79604107c2585299 Mon Sep 17 00:00:00 2001 From: Jun Tian Date: Wed, 2 Jul 2025 07:23:26 +0800 Subject: [PATCH] Revert "Upgrade orbax checkpointer to 0.11.15 (#1255)" This reverts commit 8c699b369e96d239eaff8916de4928b6ce32a75b. --- axlearn/common/checkpointer_orbax_emergency.py | 13 +++++-------- pyproject.toml | 2 +- 2 files changed, 6 insertions(+), 9 deletions(-) diff --git a/axlearn/common/checkpointer_orbax_emergency.py b/axlearn/common/checkpointer_orbax_emergency.py index f868488e0..30e97caff 100644 --- a/axlearn/common/checkpointer_orbax_emergency.py +++ b/axlearn/common/checkpointer_orbax_emergency.py @@ -751,7 +751,7 @@ def save( # including step time in total blocking time. start_t = time.perf_counter() self._get_tensor_manager(state_with_tensors).save( - step=step, args=ocp.args.Composite(state=ocp.args.PyTreeSave(item=state_with_tensors)) + step=step, args=ocp.args.PyTreeSave(item=state_with_tensors) ) time_diff = time.perf_counter() - start_t if self._composite_save_policy(step=step, evaler_summaries=self._eval_summaries): @@ -808,9 +808,7 @@ def restore( restored_state_with_tensors = tensor_manager.restore( step=step, - args=ocp.args.Composite( - state=ocp.args.PyTreeRestore(item=self._get_abstract_state(state_with_tensors)) - ), + args=ocp.args.PyTreeRestore(item=self._get_abstract_state(state_with_tensors)), ) # Merge non-tensor and tensor states by replacing leaves of the non-tensor Pytree with the # not-None leaves of the tensor Pytree. @@ -828,8 +826,7 @@ def wait_until_finished(self): self._non_tensor_manager.wait_until_finished() self._tensor_manager.wait_until_finished() - def stop(self, *, has_exception: bool = False): + def stop(self): """See `BaseCheckpointer.stop` for details.""" - self._non_tensor_manager.stop(has_exception=has_exception) - if self._tensor_manager: - self._tensor_manager.close() + self._non_tensor_manager.stop() + self._tensor_manager.close() diff --git a/pyproject.toml b/pyproject.toml index 57e415e7e..595f8990c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -154,7 +154,7 @@ mmau = [ # Orbax checkpointing. orbax = [ "humanize==4.10.0", - "orbax-checkpoint==0.11.15", + "orbax-checkpoint==0.11.1", ] # Audio dependencies. audio = [