Skip to content

Commit 9e3fc76

Browse files
committed
Add: Added checkpoint cleanup after elastic event
1 parent 065b349 commit 9e3fc76

3 files changed

Lines changed: 40 additions & 2 deletions

File tree

axlearn/common/launch_trainer.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -145,6 +145,8 @@ def get_trainer_config(
145145
)
146146
trainer_config: SpmdTrainer.Config = trainer_config_fn()
147147
trainer_config.dir = trainer_config.dir or flag_values.trainer_dir
148+
149+
print(f"Trainer Config Dir: {trainer_config.dir} by Camilo")
148150
if flag_values.mesh_selector is not None:
149151
select_mesh_config(trainer_config, mesh_selector=flag_values.mesh_selector)
150152
trainer_config.mesh_axis_names = trainer_config.mesh_axis_names or ("data", "model")

axlearn/common/launch_trainer_main.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
"""Main function for launching the trainer."""
44

55
import pathwaysutils
6+
import functools
67
from absl import app, flags
78
from pathwaysutils.elastic import elastic, manager
89

@@ -17,10 +18,10 @@
1718
def main(_):
1819
measurement.initialize(flags.FLAGS)
1920
launch.setup()
20-
# trainer_config = launch_trainer.get_trainer_config()
21+
trainer_config = launch_trainer.get_trainer_config()
2122
# trainer_config.set(recorder=config_for_function(lambda: measurement.global_recorder))
2223
# measurement.start_monitoring()
23-
24+
clean_up_checkpoints = functools.partial(utils.clean_up_checkpoints, checkpoint_dir=trainer_config.dir)
2425
if pathwaysutils.is_pathways_backend_used() and enable_elastic_training:
2526

2627
def train():
@@ -69,6 +70,7 @@ def pre_callback():
6970
max_resizes=10, # Handle up to 10 slice up or slice down transitions
7071
poll_interval=30, # Monitor thread checks inactive slice health every 30 seconds
7172
pre_callback=pre_callback,
73+
on_elastic_event_callback=clean_up_checkpoints,
7274
)(train)
7375

7476
train()

axlearn/common/utils.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
from collections.abc import Mapping, Sequence
2929
from enum import Enum
3030
from functools import cache
31+
import subprocess
3132
from typing import (
3233
Any,
3334
Callable,
@@ -109,6 +110,39 @@ def live_devices():
109110
def live_slice_indices() -> set[int]:
110111
return {d.slice_index for d in live_devices()}
111112

113+
def clean_up_checkpoints(checkpoint_dir: str):
114+
115+
print(f"Checking for incomplete checkpoint after an elastic event...Check dir: {checkpoint_dir}")
116+
117+
# 1. List the directory
118+
new_checkpoint_dir = f"{checkpoint_dir}/checkpoints/"
119+
result = subprocess.run(['gsutil', 'ls', new_checkpoint_dir], capture_output=True, text=True)
120+
121+
if result.returncode != 0:
122+
print("Failed to inspect checkpoint dir. Continuing")
123+
return
124+
print(f"Checkpoints==> {[line for line in result.stdout.splitlines()]}")
125+
checkpoints = [line for line in result.stdout.splitlines()]
126+
127+
if not checkpoints:
128+
print("Found no existing checkpoints. Continuing")
129+
return
130+
131+
# Sort naturally (Version sort) and get the last one
132+
checkpoints.sort(key=lambda x: [int(c) if c.isdigit() else c for c in re.split(r'(\d+)', x)])
133+
latest_checkpoint = checkpoints[-1]
134+
135+
print(f"Checking latest checkpoint: {latest_checkpoint}")
136+
137+
# 3. Check for commit_success file
138+
# gsutil -q stat returns 0 if found, non-zero if not
139+
stat_check = subprocess.run(['gsutil', '-q', 'stat', f"{latest_checkpoint}commit_success*"])
140+
141+
if stat_check.returncode != 0:
142+
print(f"No commit_success file found. Deleting {latest_checkpoint}...")
143+
subprocess.run(['gsutil', '-m', 'rm', '-rf', latest_checkpoint])
144+
else:
145+
print(f"Found commit_success file. Keeping {latest_checkpoint}.")
112146

113147
@dataclasses.dataclass
114148
class HybridMeshShape:

0 commit comments

Comments
 (0)