|
14 | 14 | from dags.common.vm_resource import XpkClusters |
15 | 15 | from dags.multipod.configs import gke_config |
16 | 16 | from dags.orbax.util import checkpoint_util |
| 17 | +from dags.orbax.util import test_config_util |
17 | 18 | from dags.orbax.util import validation_util |
18 | 19 | from xlml.utils.gke import zone_to_region |
19 | | -from dags.orbax.util import test_config_util |
20 | 20 |
|
21 | 21 |
|
22 | | -SCHEDULE = "45 12 * * *" if composer_env.is_prod_env() else None |
| 22 | +SCHEDULE = "15 11 * * *" if composer_env.is_prod_env() else None |
23 | 23 | DAG_TEST_NAME = "maxtext_emc_save_gcs" |
24 | 24 |
|
25 | 25 |
|
|
37 | 37 | "TPU", |
38 | 38 | "v5p-128", |
39 | 39 | ], |
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 | + ), |
41 | 44 | doc_md=""" |
42 | 45 | # Multi-tier Checkpoint Validation DAG |
43 | 46 |
|
|
101 | 104 | checkpoint_dir=test_config_util.DEFAULT_RAM_DISK, |
102 | 105 | run_name=run_name, |
103 | 106 | 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 | + ), |
105 | 110 | slice_num=slice_num, |
106 | 111 | ) |
107 | 112 |
|
|
133 | 138 | location=zone_to_region(test_config.cluster.zone), |
134 | 139 | cluster_name=test_config.cluster.name, |
135 | 140 | 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+-", |
137 | 142 | start_time=start_time, |
138 | 143 | end_time=end_time, |
139 | 144 | steps_to_validate=steps_to_validate, |
140 | 145 | ) |
141 | 146 |
|
142 | 147 | # 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 | + ) |
146 | 156 | ) |
147 | 157 |
|
148 | 158 | # Final CPC cleanup to ensure symmetric start/end |
|
151 | 161 | task_id="wait_delete_cpc_final", |
152 | 162 | )(test_config.cpc_config).as_teardown(setups=apply_cpc) |
153 | 163 |
|
| 164 | + # Airflow uses >> for task chaining, which is pointless for pylint. |
| 165 | + # pylint: disable=pointless-statement |
154 | 166 | ( |
155 | 167 | wait_delete_cpc |
156 | 168 | >> apply_cpc |
|
162 | 174 | >> validate_checkpoints_steps_gcs |
163 | 175 | >> wait_delete_cpc_final |
164 | 176 | ) |
| 177 | + # pylint: enable=pointless-statement |
0 commit comments