Skip to content

Commit 088bbd7

Browse files
fix: Adjust execution schedule for multiple DAGs (GoogleCloudPlatform#1273)
This change adjusts the execution order of Orbax DAGs to isolate a certain issue for further troubleshooting. Will change time schedule from Save -> Resume -> Restore.
1 parent 736151c commit 088bbd7

13 files changed

Lines changed: 166 additions & 57 deletions

dags/orbax/axlearn_checkpoint_regular.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,10 +27,9 @@
2727
from dags.orbax.util import test_config_util, validation_util
2828
from xlml.utils.gke import zone_to_region
2929
from xlml.utils import axlearn
30-
from xlml.apis.xpk_cluster_config import XpkClusterConfig
3130

3231

33-
SCHEDULE = "0 17 * * *" if composer_env.is_prod_env() else None
32+
SCHEDULE = "0 13 * * *" if composer_env.is_prod_env() else None
3433
DAG_TEST_NAME = "axlearn_reg_save"
3534

3635
RESERVE_TIME_FOR_OTHERS = datetime.timedelta(minutes=5)

dags/orbax/maxtext_emc_restore_gcs.py

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from xlml.utils.gke import zone_to_region
1919

2020
DAG_TEST_NAME = "maxtext_emc_orbax_res_gcs"
21-
SCHEDULE = "0 10 * * *" if composer_env.is_prod_env() else None
21+
SCHEDULE = "30 18 * * *" if composer_env.is_prod_env() else None
2222

2323
with models.DAG(
2424
dag_id=DAG_TEST_NAME,
@@ -34,7 +34,10 @@
3434
"TPU",
3535
"v5p-128",
3636
],
37-
description="DAG to verify MaxText's emergency restore from GCS checkpoints after a full cluster interruption.",
37+
description=(
38+
"DAG to verify MaxText's emergency restore from GCS"
39+
" checkpoints after a full cluster interruption."
40+
),
3841
doc_md="""
3942
# MaxText Emergency Restore from GCS Validation DAG
4043
@@ -108,7 +111,9 @@
108111
run_name=run_name,
109112
slice_num=slice_num,
110113
out_folder="maxtext_emc_orbax_res_gcs",
111-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
114+
enable_multi_tier_checkpointing=(
115+
checkpointing.enable_multi_tier_checkpointing
116+
),
112117
)
113118

114119
start_time = validation_util.generate_timestamp.override(
@@ -174,18 +179,23 @@
174179
is_local=False
175180
)
176181

177-
validate_saved_checkpoints_steps_gcs = validation_util.validate_gcs_checkpoint_files(
178-
bucket_path=(
179-
f"{test_config_util.DEFAULT_BUCKET}/{DAG_TEST_NAME}/{run_name}"
180-
),
181-
steps_to_validate=gcs_saved_steps_to_validate,
182+
validate_saved_checkpoints_steps_gcs = (
183+
validation_util.validate_gcs_checkpoint_files(
184+
bucket_path=(
185+
f"{test_config_util.DEFAULT_BUCKET}"
186+
f"/{DAG_TEST_NAME}/{run_name}"
187+
),
188+
steps_to_validate=gcs_saved_steps_to_validate,
189+
)
182190
)
183191
# Final CPC cleanup to ensure symmetric start/end
184192
wait_delete_cpc_final = checkpoint_util.wait_for_cpc_deletion.override(
185193
trigger_rule=TriggerRule.ALL_DONE,
186194
task_id="wait_delete_cpc_final",
187195
)(test_config.cpc_config).as_teardown(setups=apply_cpc)
188196

197+
# Airflow uses >> for task chaining, which is pointless for pylint.
198+
# pylint: disable=pointless-statement
189199
(
190200
wait_delete_cpc
191201
>> apply_cpc
@@ -198,3 +208,4 @@
198208
>> validate_saved_checkpoints_steps_gcs
199209
>> wait_delete_cpc_final
200210
)
211+
# pylint: enable=pointless-statement

dags/orbax/maxtext_emc_restore_local.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
from xlml.utils.gke import zone_to_region
1616

1717
DAG_TEST_NAME = "maxtext_emc_orbax_res_local"
18-
SCHEDULE = "45 10 * * *" if composer_env.is_prod_env() else None
18+
SCHEDULE = "15 19 * * *" if composer_env.is_prod_env() else None
1919

2020

2121
with models.DAG(
@@ -32,7 +32,10 @@
3232
"TPU",
3333
"v5p-128",
3434
],
35-
description="DAG to verify MaxText's emergency restore from local checkpoints after a node interruption.",
35+
description=(
36+
"DAG to verify MaxText's emergency restore from local"
37+
" checkpoints after a node interruption."
38+
),
3639
doc_md="""
3740
# MaxText Emergency Restore from Local Checkpoint Validation DAG
3841
@@ -106,7 +109,9 @@
106109
run_name=run_name,
107110
slice_num=slice_num,
108111
out_folder="maxtext_emc_orbax_res_local",
109-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
112+
enable_multi_tier_checkpointing=(
113+
checkpointing.enable_multi_tier_checkpointing
114+
),
110115
)
111116

112117
start_time = validation_util.generate_timestamp.override(
@@ -183,6 +188,8 @@
183188
task_id="wait_delete_cpc_final",
184189
)(test_config.cpc_config).as_teardown(setups=apply_cpc)
185190

191+
# Airflow uses >> for task chaining, which is pointless for pylint.
192+
# pylint: disable=pointless-statement
186193
(
187194
wait_delete_cpc
188195
>> apply_cpc
@@ -195,3 +202,4 @@
195202
>> validate_log
196203
>> wait_delete_cpc_final
197204
)
205+
# pylint: enable=pointless-statement

dags/orbax/maxtext_emc_resume_gcs.py

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
from dags.orbax.util import checkpoint_util
2222
from xlml.utils.gke import zone_to_region
2323

24-
SCHEDULE = "30 15 * * *" if composer_env.is_prod_env() else None
24+
SCHEDULE = "15 14 * * *" if composer_env.is_prod_env() else None
2525
DAG_TEST_NAME = "maxtext_emc_resume_from_gcs"
2626

2727
with models.DAG(
@@ -38,7 +38,10 @@
3838
"TPU",
3939
"v5p-128",
4040
],
41-
description="A DAG to test MaxText Emergency Checkpoint Manager GCS restore functionality.",
41+
description=(
42+
"A DAG to test MaxText Emergency Checkpoint Manager"
43+
" GCS restore functionality."
44+
),
4245
doc_md="""
4346
# Emergency Checkpoint Manager GCS Restore Validation DAG
4447
@@ -136,7 +139,9 @@
136139
run_name=run_name,
137140
slice_num=slice_num,
138141
out_folder=out_folder,
139-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
142+
enable_multi_tier_checkpointing=(
143+
checkpointing.enable_multi_tier_checkpointing
144+
),
140145
)
141146

142147
start_time = validation_util.generate_timestamp.override(
@@ -169,7 +174,9 @@
169174
run_name=run_name,
170175
slice_num=slice_num,
171176
out_folder=out_folder,
172-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
177+
enable_multi_tier_checkpointing=(
178+
checkpointing.enable_multi_tier_checkpointing
179+
),
173180
)
174181

175182
resume_training_run = gke_config.get_gke_config(
@@ -249,6 +256,8 @@
249256
task_id="wait_delete_cpc_final",
250257
)(test_config.cpc_config).as_teardown(setups=apply_first_cpc)
251258

259+
# Airflow uses >> for task chaining, which is pointless for pylint.
260+
# pylint: disable=pointless-statement
252261
(
253262
wait_delete_cpc
254263
>> apply_first_cpc
@@ -264,3 +273,4 @@
264273
>> validate_saved_checkpoints_steps_gcs
265274
>> wait_delete_cpc_final
266275
)
276+
# pylint: enable=pointless-statement

dags/orbax/maxtext_emc_save_gcs.py

Lines changed: 21 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -14,12 +14,12 @@
1414
from dags.common.vm_resource import XpkClusters
1515
from dags.multipod.configs import gke_config
1616
from dags.orbax.util import checkpoint_util
17+
from dags.orbax.util import test_config_util
1718
from dags.orbax.util import validation_util
1819
from xlml.utils.gke import zone_to_region
19-
from dags.orbax.util import test_config_util
2020

2121

22-
SCHEDULE = "45 12 * * *" if composer_env.is_prod_env() else None
22+
SCHEDULE = "15 11 * * *" if composer_env.is_prod_env() else None
2323
DAG_TEST_NAME = "maxtext_emc_save_gcs"
2424

2525

@@ -37,7 +37,10 @@
3737
"TPU",
3838
"v5p-128",
3939
],
40-
description="DAG that verifies the orbax multi-tier checkpointing saving functionality with replicator to GCS bucket",
40+
description=(
41+
"DAG that verifies the orbax multi-tier checkpointing saving"
42+
" functionality with replicator to GCS bucket"
43+
),
4144
doc_md="""
4245
# Multi-tier Checkpoint Validation DAG
4346
@@ -101,7 +104,9 @@
101104
checkpoint_dir=test_config_util.DEFAULT_RAM_DISK,
102105
run_name=run_name,
103106
out_folder="maxtext_emc_orbax_save_gcs",
104-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
107+
enable_multi_tier_checkpointing=(
108+
checkpointing.enable_multi_tier_checkpointing
109+
),
105110
slice_num=slice_num,
106111
)
107112

@@ -133,16 +138,21 @@
133138
location=zone_to_region(test_config.cluster.zone),
134139
cluster_name=test_config.cluster.name,
135140
ram_disk="gcs",
136-
pod_pattern=f"{test_config.short_id}-emc.*-0-\d+-",
141+
pod_pattern=rf"{test_config.short_id}-emc.*-0-\d+-",
137142
start_time=start_time,
138143
end_time=end_time,
139144
steps_to_validate=steps_to_validate,
140145
)
141146

142147
# Validate that GCS restore happened during the second training run
143-
validate_checkpoints_steps_gcs = validation_util.validate_gcs_checkpoint_files(
144-
bucket_path=f"{test_config_util.DEFAULT_BUCKET}/maxtext_emc_orbax_save_gcs/{run_name}",
145-
steps_to_validate=steps_to_validate,
148+
validate_checkpoints_steps_gcs = (
149+
validation_util.validate_gcs_checkpoint_files(
150+
bucket_path=(
151+
f"{test_config_util.DEFAULT_BUCKET}"
152+
f"/maxtext_emc_orbax_save_gcs/{run_name}"
153+
),
154+
steps_to_validate=steps_to_validate,
155+
)
146156
)
147157

148158
# Final CPC cleanup to ensure symmetric start/end
@@ -151,6 +161,8 @@
151161
task_id="wait_delete_cpc_final",
152162
)(test_config.cpc_config).as_teardown(setups=apply_cpc)
153163

164+
# Airflow uses >> for task chaining, which is pointless for pylint.
165+
# pylint: disable=pointless-statement
154166
(
155167
wait_delete_cpc
156168
>> apply_cpc
@@ -162,3 +174,4 @@
162174
>> validate_checkpoints_steps_gcs
163175
>> wait_delete_cpc_final
164176
)
177+
# pylint: enable=pointless-statement

dags/orbax/maxtext_mtc_emergency_save_local.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525
from xlml.utils.gke import zone_to_region
2626

2727

28-
SCHEDULE = "0 12 * * *" if composer_env.is_prod_env() else None
28+
SCHEDULE = "45 11 * * *" if composer_env.is_prod_env() else None
2929
DAG_TEST_NAME = "maxtext_emc_and_mtc_orbax_save_local"
3030

3131

dags/orbax/maxtext_mtc_restore_local.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from xlml.utils.gke import zone_to_region
1919

2020
DAG_TEST_NAME = "maxtext_mtc_orbax_res_local"
21-
SCHEDULE = "45 13 * * *" if composer_env.is_prod_env() else None
21+
SCHEDULE = "30 21 * * *" if composer_env.is_prod_env() else None
2222

2323
with models.DAG(
2424
dag_id=DAG_TEST_NAME,
@@ -35,7 +35,10 @@
3535
"TPU",
3636
"v5p-128",
3737
],
38-
description="DAG to verify MaxText's multi-tier restore from local checkpoints after a node interruption.",
38+
description=(
39+
"DAG to verify MaxText's multi-tier restore from local"
40+
" checkpoints after a node interruption."
41+
),
3942
doc_md="""
4043
# MaxText Multi-tier Restore from Local Checkpoint Validation DAG
4144
@@ -109,7 +112,9 @@
109112
run_name=run_name,
110113
slice_num=slice_num,
111114
out_folder="maxtext_mtc_orbax_res_local",
112-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
115+
enable_multi_tier_checkpointing=(
116+
checkpointing.enable_multi_tier_checkpointing
117+
),
113118
)
114119

115120
start_time = validation_util.generate_timestamp.override(
@@ -188,6 +193,8 @@
188193
task_id="wait_delete_cpc_final",
189194
)(test_config.cpc_config).as_teardown(setups=apply_cpc)
190195

196+
# Airflow uses >> for task chaining, which is pointless for pylint.
197+
# pylint: disable=pointless-statement
191198
(
192199
wait_delete_cpc
193200
>> apply_cpc
@@ -200,3 +207,4 @@
200207
>> validate_local_saved_steps
201208
>> wait_delete_cpc_final
202209
)
210+
# pylint: enable=pointless-statement

dags/orbax/maxtext_mtc_resume_gcs.py

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
from dags.orbax.util import checkpoint_util
2222
from xlml.utils.gke import zone_to_region
2323

24-
SCHEDULE = "45 18 * * *" if composer_env.is_prod_env() else None
24+
SCHEDULE = "45 15 * * *" if composer_env.is_prod_env() else None
2525
DAG_TEST_NAME = "maxtext_mtc_resume_from_gcs"
2626

2727
with models.DAG(
@@ -39,7 +39,10 @@
3939
"TPU",
4040
"v5p-128",
4141
],
42-
description="A DAG to test MaxText Multi-tier Checkpointing (MTC) GCS restore functionality.",
42+
description=(
43+
"A DAG to test MaxText Multi-tier Checkpointing (MTC)"
44+
" GCS restore functionality."
45+
),
4346
doc_md="""
4447
# Multi-tier Checkpointing (MTC) GCS Restore Validation DAG
4548
@@ -139,7 +142,9 @@
139142
run_name=run_name,
140143
slice_num=slice_num,
141144
out_folder=out_folder,
142-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
145+
enable_multi_tier_checkpointing=(
146+
checkpointing.enable_multi_tier_checkpointing
147+
),
143148
)
144149

145150
start_time = validation_util.generate_timestamp.override(
@@ -175,7 +180,9 @@
175180
run_name=run_name,
176181
slice_num=slice_num,
177182
out_folder=out_folder,
178-
enable_multi_tier_checkpointing=checkpointing.enable_multi_tier_checkpointing,
183+
enable_multi_tier_checkpointing=(
184+
checkpointing.enable_multi_tier_checkpointing
185+
),
179186
)
180187

181188
resume_training_run = gke_config.get_gke_config(
@@ -239,7 +246,8 @@
239246
)
240247
)
241248

242-
# Validate that MTC checkpoint files exist in GCS bucket with correct backup folder structure
249+
# Validate that MTC checkpoint files exist in GCS bucket with
250+
# correct backup folder structure
243251
validate_mtc_gcs_files = (
244252
validation_util.validate_gcs_checkpoint_files.override(
245253
task_id="validate_mtc_gcs_files"
@@ -256,6 +264,8 @@
256264
task_id="wait_delete_cpc_final",
257265
)(test_config.cpc_config).as_teardown(setups=apply_cpc)
258266

267+
# Airflow uses >> for task chaining, which is pointless for pylint.
268+
# pylint: disable=pointless-statement
259269
(
260270
wait_delete_cpc
261271
>> apply_cpc
@@ -272,3 +282,4 @@
272282
>> validate_mtc_gcs_files
273283
>> wait_delete_cpc_final
274284
)
285+
# pylint: enable=pointless-statement

0 commit comments

Comments
 (0)