(
self,
ckpt_paths: Optional[Sequence[str]],
expect_restore_step: int,
expect_saved_steps: Sequence[int],
listdir_add_trailing_slash: bool,
)
| 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 |
nothing calls this directly
no test coverage detected