MCPcopy Create free account
hub / github.com/apple/axlearn / test_garbage_collection

Method test_garbage_collection

axlearn/common/checkpointer_test.py:759–834  ·  view source on GitHub ↗
(
        self,
        ckpt_paths: Optional[Sequence[str]],
        expect_restore_step: int,
        expect_saved_steps: Sequence[int],
        listdir_add_trailing_slash: bool,
    )

Source from the content-addressed store, hash-verified

757 ),
758 ),
759 listdir_add_trailing_slash=[True, False],
760 )
761 def test_garbage_collection(
762 self,
763 ckpt_paths: Optional[Sequence[str]],
764 expect_restore_step: int,
765 expect_saved_steps: Sequence[int],
766 listdir_add_trailing_slash: bool,
767 ):
768 mesh_shape = (1, 1)
769 if not test_utils.is_supported_mesh_shape(mesh_shape):
770 return
771
772 orig_listdir = listdir
773
774 def patch_tf_io_behavior(*args):
775 out = orig_listdir(*args)
776 return [x + "/" for x in out if not x.endswith("/")]
777
778 # pylint: disable=line-too-long
779 with (
780 _mesh(mesh_shape),
781 (
782 mock.patch("axlearn.common.file_system.listdir", patch_tf_io_behavior)
783 if listdir_add_trailing_slash
784 else nullcontext()
785 ),
786 tempfile.TemporaryDirectory() as temp_dir,
787 ):
788 cfg = Checkpointer.default_config().set(
789 name="test",
790 dir=temp_dir,
791 keep_last_n=3,
792 keep_every_n_steps=2,
793 gc_loop_interval_seconds=1,
794 )
795 cfg.save_policy.min_step = 0
796
797 # Running gc for non-existent dir shouldn't fail.
798 ckpt_fake = cfg.clone(dir=os.path.join(temp_dir, "fake_dir")).instantiate(parent=None)
799 ckpt_fake._run_garbage_collection()
800
801 ckpt: Checkpointer = cfg.instantiate(parent=None)
802 state = dict(x=jnp.zeros([], dtype=jnp.int32))
803
804 for step in range(10):
805 ckpt.save(step=step, state=state)
806 ckpt.wait_until_finished()
807
808 # Mock out the checkpoints that are committed.
809 if ckpt_paths:
810 ckpt_paths = [os.path.join(temp_dir, f"step_{i:08d}") for i in ckpt_paths]
811 else:
812 ckpt_paths = Checkpointer.checkpoint_paths(cfg.dir)
813
814 with mock.patch.object(Checkpointer, "checkpoint_paths", return_value=ckpt_paths):
815 ckpt._run_garbage_collection()
816

Callers

nothing calls this directly

Calls 14

cloneMethod · 0.80
joinMethod · 0.80
assertNestedEqualMethod · 0.80
_meshFunction · 0.70
patchMethod · 0.45
setMethod · 0.45
default_configMethod · 0.45
instantiateMethod · 0.45
saveMethod · 0.45
wait_until_finishedMethod · 0.45
checkpoint_pathsMethod · 0.45

Tested by

no test coverage detected