Skip to content

Commit 30513ab

Browse files
authored
Merge branch 'res02-restore-gcs' into sav01-save-local
2 parents ef732da + a7a6949 commit 30513ab

7 files changed

Lines changed: 601 additions & 0 deletions

File tree

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,122 @@
1+
"""Add commentMore actions
2+
A DAG to run MaxText multi-tier checkpointing tests (phase2: save & validate).
3+
"""
4+
5+
import datetime
6+
from airflow import models
7+
from dags import composer_env, gcs_bucket
8+
from dags.common import test_owner
9+
from dags.common.vm_resource import DockerImage, XpkClusters
10+
from dags.multipod.configs import gke_config
11+
from dags.multipod.configs.common import SetupMode
12+
from xlml.utils import log_explorer
13+
from xlml.utils import xpk
14+
15+
SCHEDULE = "0 10 * * *" if composer_env.is_prod_env() else None
16+
17+
with models.DAG(
18+
dag_id="maxtext_multi_tier_sav01_save_local",
19+
schedule_interval=SCHEDULE,
20+
tags=[
21+
"multipod_team",
22+
"maxtext",
23+
"multi_tier_p2_chkpt_save_local",
24+
"nightly",
25+
],
26+
start_date=datetime.datetime(2025, 5, 22),
27+
catchup=False,
28+
concurrency=2,
29+
) as dag:
30+
base_output_directory = (
31+
f"{gcs_bucket.BASE_OUTPUT_DIR}/maxtext_multi_tier_sav01_save_local"
32+
)
33+
docker_images = [(
34+
SetupMode.JAX_STABLE_STACK,
35+
DockerImage.MAXTEXT_TPU_JAX_NIGHTLY,
36+
)]
37+
ram_disk = "/local"
38+
test_configs = {"v5p-8": [2]}
39+
clusters = {"v5p-8": XpkClusters.TPU_V5P_8_CLUSTER}
40+
step = 100
41+
local_checkpoint_period = 20
42+
replicator_backup_interval_minutes = 1
43+
use_replicator = "True"
44+
name_prefix = "maxtext_phase2_chkpt_save"
45+
46+
for mode, image in docker_images:
47+
for accelerator, slices in test_configs.items():
48+
for slice_num in slices:
49+
run_time = datetime.datetime.now().strftime("%Y-%m-%d-%H")
50+
run_name = f"{name_prefix}-{slice_num}x-{accelerator}_{run_time}"
51+
workload_command = (
52+
"export TPU_PREMAPPED_BUFFER_SIZE=52428800000 && "
53+
"export TPU_PREMAPPED_BUFFER_TRANSFER_THRESHOLD_BYTES=52428800000 && "
54+
"python3 -m MaxText.train MaxText/configs/base.yml remat_policy=full "
55+
f"global_parameter_scale=1 base_output_directory={base_output_directory} "
56+
f"dataset_type=synthetic steps={step} per_device_batch_size=1 "
57+
"max_target_length=256 "
58+
"reuse_example_batch=1 enable_emergency_checkpoint=true "
59+
f"local_checkpoint_directory={ram_disk} local_checkpoint_period={local_checkpoint_period} "
60+
f"use_replicator_service={use_replicator} replicator_backup_interval_minutes={replicator_backup_interval_minutes} "
61+
f"run_name={run_name}",
62+
)
63+
64+
start_time = xpk.generate_timestamp()
65+
66+
# make launch test_name unique
67+
maxtext_phase2_chkpt_test = gke_config.get_gke_config(
68+
num_slices=slice_num,
69+
cluster=clusters[accelerator],
70+
time_out_in_min=60,
71+
test_name=f"maxtext_phase2_chkpt_save",
72+
run_model_cmds=workload_command,
73+
docker_image=image.value,
74+
test_owner=test_owner.ERNIE_C,
75+
).run(
76+
ramdisk_directory=ram_disk,
77+
mtc_enabled=True,
78+
xpk_branch="main",
79+
skip_post_process=True,
80+
)
81+
82+
# cleanup run: unique test_name
83+
cleanup_command = (f"rm -rf {ram_disk}/*",)
84+
ram_disk_cleanup = gke_config.get_gke_config(
85+
num_slices=slice_num,
86+
cluster=clusters[accelerator],
87+
time_out_in_min=60,
88+
test_name=f"maxtext_phase2_chkpt_test-cleanup",
89+
run_model_cmds=cleanup_command,
90+
docker_image=image.value,
91+
test_owner=test_owner.ERNIE_C,
92+
).run(
93+
ramdisk_directory=ram_disk,
94+
mtc_enabled=True,
95+
xpk_branch="main",
96+
skip_post_process=True,
97+
)
98+
99+
end_time = xpk.generate_timestamp()
100+
vali_step = step - 1
101+
vali_step_list = [
102+
i for i in range(0, vali_step, local_checkpoint_period)
103+
]
104+
vali_step_list.append(vali_step)
105+
106+
validate_log = log_explorer.validate_log_with_step(
107+
project_id=clusters[accelerator].project,
108+
location=clusters[accelerator].zone[:-2],
109+
cluster_name=clusters[accelerator].name,
110+
text_filter="Finished asynchronous save `(blocking` `+` `background)` in",
111+
start_time=start_time,
112+
end_time=end_time,
113+
vali_step_list=vali_step_list,
114+
)
115+
116+
(
117+
start_time
118+
>> maxtext_phase2_chkpt_test
119+
>> ram_disk_cleanup
120+
>> end_time
121+
>> validate_log
122+
)

dags/orbax/maxtext_multi_tier_chechpoint_save_local.py

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,26 +23,44 @@
2323
"multi_tier_p2_chkpt_save_local",
2424
"nightly",
2525
],
26+
<<<<<<< sav01-save-local
2627
start_date=datetime.datetime(2025, 6, 12),
28+
=======
29+
start_date=datetime.datetime(2025, 6, 27),
30+
>>>>>>> res02-restore-gcs
2731
catchup=False,
2832
concurrency=2,
2933
) as dag:
3034
base_output_directory = (
3135
f"{gcs_bucket.MTC_AUTOMATION_BUCKET}/maxtext_multi_tier_sav01_save_local"
3236
)
37+
<<<<<<< sav01-save-local
3338
docker_images = [(
3439
SetupMode.NIGHTLY,
3540
DockerImage.MAXTEXT_TPU_JAX_NIGHTLY,
3641
)]
42+
=======
43+
docker_images = [
44+
(
45+
SetupMode.NIGHTLY,
46+
DockerImage.MAXTEXT_TPU_JAX_NIGHTLY,
47+
)
48+
]
49+
>>>>>>> res02-restore-gcs
3750
ram_disk = "/local"
3851
test_configs = {"v5p-128": [2]}
3952
clusters = {"v5p-128": XpkClusters.TPU_V5P_128_CLUSTER}
4053
step = 100
4154
local_checkpoint_period = 20
4255
replicator_backup_interval_minutes = 1
4356
use_replicator = "True"
57+
<<<<<<< sav01-save-local
4458
name_prefix = "maxtext_phase2_chkpt_save"
4559
model_name = "llama2-7b"
60+
=======
61+
model_name = "llama2-7b"
62+
name_prefix = f"maxtext_{model_name}_chkpt_save"
63+
>>>>>>> res02-restore-gcs
4664

4765
for mode, image in docker_images:
4866
for accelerator, slices in test_configs.items():
Lines changed: 216 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,216 @@
1+
"""A DAG to run MaxText multi-tier checkpointing tests (phase2: restore & validate)."""
2+
3+
import datetime
4+
from airflow import models
5+
from dags import composer_env, gcs_bucket
6+
from dags.common import test_owner
7+
from dags.common.vm_resource import DockerImage, XpkClusters
8+
from dags.multipod.configs import gke_config
9+
from dags.multipod.configs.common import SetupMode
10+
from xlml.utils import log_explorer
11+
from xlml.utils import orbax
12+
from xlml.utils import xpk
13+
14+
SCHEDULE = "0 10 * * *" if composer_env.is_prod_env() else None
15+
16+
with models.DAG(
17+
dag_id="maxtext_multi_tier_res02_restore_gcs",
18+
schedule_interval=SCHEDULE,
19+
tags=[
20+
"multipod_team",
21+
"maxtext",
22+
"multi_tier_chkpt_restore_gcs",
23+
"nightly",
24+
],
25+
start_date=datetime.datetime(2025, 6, 25),
26+
catchup=False,
27+
concurrency=2,
28+
) as dag:
29+
base_output_directory = (
30+
f"{gcs_bucket.BASE_OUTPUT_DIR}/maxtext_multi_tier_res02_restore_gcs"
31+
)
32+
docker_images = [
33+
(
34+
SetupMode.NIGHTLY,
35+
DockerImage.MAXTEXT_TPU_JAX_NIGHTLY,
36+
)
37+
]
38+
ram_disk = "/local"
39+
test_configs = {"v5p-128": [2]}
40+
clusters = {"v5p-128": XpkClusters.TPU_V5P_128_CLUSTER}
41+
step = 200
42+
restore_step = 300
43+
local_checkpoint_period = 20
44+
replicator_backup_interval_minutes = 1
45+
use_replicator = "True"
46+
model_name = "llama2-7b"
47+
name_prefix = f"maxtext_{model_name}_chkpt_restore"
48+
49+
for mode, image in docker_images:
50+
for accelerator, slices in test_configs.items():
51+
for slice_num in slices:
52+
cpc = (
53+
clusters[accelerator].project,
54+
clusters[accelerator].zone[:-2],
55+
clusters[accelerator].name,
56+
gcs_bucket.MTC_AUTOMATION_BUCKET.split("gs://")[1],
57+
"ct5p-hightpu-4t",
58+
"google.com/tpu",
59+
"800000Mi",
60+
)
61+
delete_cpc = orbax.delete_cpc(*cpc)
62+
apply_cpc = orbax.apply_cpc(*cpc)
63+
run_time = datetime.datetime.now().strftime("%Y-%m-%d-%H")
64+
run_name = f"{name_prefix}-{slice_num}x-{accelerator}_{run_time}"
65+
bucket_name = f"{gcs_bucket.BASE_OUTPUT_DIR}/{run_name}"
66+
workload_command = (
67+
"export TPU_PREMAPPED_BUFFER_SIZE=52428800000 && "
68+
"export TPU_PREMAPPED_BUFFER_TRANSFER_THRESHOLD_BYTES=52428800000 && "
69+
"python3 -m MaxText.train MaxText/configs/base.yml remat_policy=full "
70+
f"global_parameter_scale=1 base_output_directory={base_output_directory} "
71+
f"dataset_type=synthetic steps={step} per_device_batch_size=1 "
72+
"max_target_length=256 "
73+
"reuse_example_batch=1 enable_emergency_checkpoint=true "
74+
f"local_checkpoint_directory={ram_disk} local_checkpoint_period={local_checkpoint_period} "
75+
f"use_replicator_service={use_replicator} replicator_backup_interval_minutes={replicator_backup_interval_minutes} "
76+
f"run_name={run_name}",
77+
)
78+
workload_command_restore = (
79+
"export TPU_PREMAPPED_BUFFER_SIZE=52428800000 && "
80+
"export TPU_PREMAPPED_BUFFER_TRANSFER_THRESHOLD_BYTES=52428800000 && "
81+
"python3 -m MaxText.train MaxText/configs/base.yml remat_policy=full "
82+
f"global_parameter_scale=1 base_output_directory={base_output_directory} "
83+
f"dataset_type=synthetic steps={restore_step} per_device_batch_size=1 "
84+
"max_target_length=256 "
85+
"reuse_example_batch=1 enable_emergency_checkpoint=true "
86+
f"local_checkpoint_directory={ram_disk} local_checkpoint_period={local_checkpoint_period} "
87+
f"use_replicator_service={use_replicator} replicator_backup_interval_minutes={replicator_backup_interval_minutes} "
88+
f"run_name={run_name}",
89+
)
90+
91+
workload_id = xpk.generate_workload_id(f"{run_name}")
92+
93+
start_time = log_explorer.generate_timestamp()
94+
95+
# make launch test_name unique
96+
maxtext_phase2_chkpt_test = gke_config.get_gke_config(
97+
num_slices=slice_num,
98+
cluster=clusters[accelerator],
99+
time_out_in_min=60,
100+
test_name=f"{name_prefix}",
101+
run_model_cmds=workload_command,
102+
docker_image=image.value,
103+
test_owner=test_owner.ERNIE_C,
104+
).run_with_workload_id(
105+
ramdisk_directory=ram_disk,
106+
mtc_enabled=True,
107+
xpk_branch="main",
108+
skip_post_process=True,
109+
workload_id=workload_id,
110+
)
111+
112+
# cleanup run: unique test_name
113+
cleanup_command = (f"rm -rf {ram_disk}/*",)
114+
ram_disk_cleanup = gke_config.get_gke_config(
115+
num_slices=slice_num,
116+
cluster=clusters[accelerator],
117+
time_out_in_min=60,
118+
test_name=f"{name_prefix}-cleanup",
119+
run_model_cmds=cleanup_command,
120+
docker_image=image.value,
121+
test_owner=test_owner.ERNIE_C,
122+
).run(
123+
ramdisk_directory=ram_disk,
124+
mtc_enabled=True,
125+
xpk_branch="main",
126+
skip_post_process=True,
127+
)
128+
129+
end_time = log_explorer.generate_timestamp()
130+
validate_gcs_bucket_save_step = log_explorer.validate_log_with_gcs(
131+
project_id=clusters[accelerator].project,
132+
location=clusters[accelerator].zone[:-2],
133+
cluster_name=clusters[accelerator].name,
134+
text_filter="Successful: backup for step",
135+
namespace="gke-managed-checkpointing",
136+
container_name="replication-worker",
137+
pod_pattern="multitier-driver",
138+
start_time=start_time,
139+
end_time=end_time,
140+
bucket_name=bucket_name,
141+
)
142+
143+
restore_start_time = log_explorer.generate_timestamp()
144+
145+
maxtext_phase2_chkpt_restore = gke_config.get_gke_config(
146+
num_slices=slice_num,
147+
cluster=clusters[accelerator],
148+
time_out_in_min=60,
149+
test_name=f"{name_prefix}_restore",
150+
run_model_cmds=workload_command_restore,
151+
docker_image=image.value,
152+
test_owner=test_owner.ERNIE_C,
153+
).run_with_workload_id(
154+
ramdisk_directory=ram_disk,
155+
mtc_enabled=True,
156+
xpk_branch="main",
157+
skip_post_process=True,
158+
workload_id=workload_id,
159+
)
160+
161+
# cleanup run: unique test_name
162+
cleanup_command = (f"rm -rf {ram_disk}/*",)
163+
ram_disk_cleanup_restore = gke_config.get_gke_config(
164+
num_slices=slice_num,
165+
cluster=clusters[accelerator],
166+
time_out_in_min=60,
167+
test_name=f"{name_prefix}-cleanup2",
168+
run_model_cmds=cleanup_command,
169+
docker_image=image.value,
170+
test_owner=test_owner.ERNIE_C,
171+
).run(
172+
ramdisk_directory=ram_disk,
173+
mtc_enabled=True,
174+
xpk_branch="main",
175+
skip_post_process=True,
176+
)
177+
178+
restore_end_time = log_explorer.generate_timestamp()
179+
180+
validate_gcs_bucket_restore_step = log_explorer.validate_log_exist(
181+
project_id=clusters[accelerator].project,
182+
location=clusters[accelerator].zone[:-2],
183+
cluster_name=clusters[accelerator].name,
184+
text_filter=f"Restoring from backup checkpoint {validate_gcs_bucket_save_step}",
185+
namespace="gke-managed-checkpointing",
186+
container_name="replication-worker",
187+
pod_pattern="multitier-driver",
188+
start_time=restore_start_time,
189+
end_time=restore_end_time,
190+
)
191+
192+
validate_gcs_bucket_restore_file = log_explorer.validate_log_exist(
193+
project_id=clusters[accelerator].project,
194+
location=clusters[accelerator].zone[:-2],
195+
cluster_name=clusters[accelerator].name,
196+
text_filter="copy backup/gcs/ to local/client/",
197+
namespace="gke-managed-checkpointing",
198+
container_name="replication-worker",
199+
pod_pattern="multitier-driver",
200+
start_time=restore_start_time,
201+
end_time=restore_end_time,
202+
)
203+
204+
(
205+
start_time
206+
>> maxtext_phase2_chkpt_test
207+
>> ram_disk_cleanup
208+
>> end_time
209+
>> validate_gcs_bucket_save_step
210+
>> restore_start_time
211+
>> maxtext_phase2_chkpt_restore
212+
>> ram_disk_cleanup_restore
213+
>> restore_end_time
214+
>> validate_gcs_bucket_restore_step
215+
>> validate_gcs_bucket_restore_file
216+
)

0 commit comments

Comments
 (0)