Skip to content

Commit bd7ea0f

Browse files
authored
feat: Implement GCS-based DAG Configuration with Airflow TaskFlow (GoogleCloudPlatform#1097)
This change refactors the configuration management for several TPU observability DAGs to use a centralized YAML file stored in Google Cloud Storage. This approach replaces hardcoded parameters and cumbersome Airflow Variables, enabling more flexible and manageable DAG configurations.
1 parent 8d46814 commit bd7ea0f

9 files changed

Lines changed: 269 additions & 166 deletions

dags/tpu_observability/configs/common.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,3 +19,8 @@ class MachineConfigMap(enum.Enum):
1919
tpu_topology="4x4",
2020
machine_version=MachineVersion.CT6E_STAND_4T,
2121
)
22+
23+
24+
GCS_CONFIG_PATH = (
25+
"gs://ml-auto-solutions-dag-configs/tpu_observability/dag_config.yaml"
26+
)

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 13 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -21,11 +21,10 @@
2121
from airflow.utils.task_group import TaskGroup
2222

2323
from dags import composer_env
24-
from dags.common.vm_resource import Region, Zone
2524
from dags.tpu_observability.utils import jobset_util as jobset
2625
from dags.tpu_observability.utils import node_pool_util as node_pool
2726
from dags.tpu_observability.utils.jobset_util import JobSet, Workload
28-
from dags.tpu_observability.configs.common import MachineConfigMap
27+
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
2928

3029

3130
# Keyword arguments are generated dynamically at runtime (pylint does not
@@ -70,32 +69,8 @@
7069
timeout, and fail.
7170
""",
7271
) as dag:
73-
cluster_name = "tpu-observability-automation"
74-
cluster_name += "-prod" if composer_env.is_prod_env() else "-dev"
75-
7672
for machine in MachineConfigMap:
7773
config = machine.value
78-
cluster_info = node_pool.Info(
79-
project_id=models.Variable.get("PROJECT_ID", default_var="cienet-cmcs"),
80-
cluster_name=models.Variable.get(
81-
"CLUSTER_NAME", default_var=cluster_name
82-
),
83-
node_pool_name=models.Variable.get(
84-
"NODE_POOL_NAME", default_var="jobset-ttr-rollback-v6e"
85-
),
86-
region=models.Variable.get(
87-
"REGION", default_var=Region.US_CENTRAL1.value
88-
),
89-
location=models.Variable.get(
90-
"LOCATION", default_var=Region.US_CENTRAL1.value
91-
),
92-
node_locations=models.Variable.get(
93-
"LOCATIONS", default_var=Zone.US_CENTRAL1_B.value
94-
),
95-
num_nodes=models.Variable.get("NUM_NODES", default_var=4),
96-
machine_type=config.machine_version.value,
97-
tpu_topology=config.tpu_topology,
98-
)
9974

10075
jobset_config = JobSet(
10176
jobset_name="ttr-rollback-v6e-workload",
@@ -118,9 +93,18 @@
11893
with TaskGroup( # pylint: disable=unexpected-keyword-arg
11994
group_id=f"v{config.tpu_version.value}"
12095
):
96+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
97+
task_id="build_node_pool_info_from_gcs_yaml"
98+
)(
99+
gcs_path=GCS_CONFIG_PATH,
100+
dag_name="jobset_rollback_ttr",
101+
is_prod=composer_env.is_prod_env(),
102+
machine_type=config.machine_version.value,
103+
tpu_topology=config.tpu_topology,
104+
)
105+
121106
create_node_pool = node_pool.create(
122107
node_pool=cluster_info,
123-
reservation="cloudtpu-20251107233000-1246578561",
124108
)
125109

126110
start_workload = jobset.run_workload(
@@ -161,7 +145,8 @@
161145
# Airflow uses >> for task chaining, which is pointless for pylint.
162146
# pylint: disable=pointless-statement
163147
(
164-
create_node_pool
148+
cluster_info
149+
>> create_node_pool
165150
>> start_workload
166151
>> ensure_all_pods_running
167152
>> rollback_node_pool

dags/tpu_observability/multi_host_nodepool_rollback_dag.py

Lines changed: 15 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -20,16 +20,14 @@
2020
import datetime
2121

2222
from airflow import models
23-
from airflow.models import Variable
2423
from airflow.utils.task_group import TaskGroup
2524
from airflow.utils.trigger_rule import TriggerRule
2625

2726
from dags import composer_env
2827
from dags.common import test_owner
2928
from dags.map_reproducibility.utils import constants
30-
from dags.common.vm_resource import Region, Zone
29+
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
3130
from dags.tpu_observability.utils import node_pool_util as node_pool
32-
from dags.tpu_observability.configs.common import MachineConfigMap
3331

3432

3533
# Keyword arguments are generated dynamically at runtime (pylint does not
@@ -74,31 +72,27 @@
7472
) as dag:
7573
for machine in MachineConfigMap:
7674
config = machine.value
77-
cluster_name = "tpu-observability-automation"
78-
cluster_name += "-prod" if composer_env.is_prod_env() else "-dev"
79-
node_pool_info = node_pool.Info(
80-
project_id="cienet-cmcs",
81-
cluster_name=cluster_name,
82-
node_pool_name=Variable.get(
83-
"NODE_POOL_NAME", default_var="multi-host-nodepool-rollback-auto"
84-
),
85-
location=Variable.get("LOCATION", default_var=Region.US_CENTRAL1.value),
86-
node_locations=Variable.get(
87-
"NODE_LOCATIONS", default_var=Zone.US_CENTRAL1_B.value
88-
),
89-
num_nodes=Variable.get("NUM_NODES", default_var=4),
90-
machine_type=config.machine_version.value,
91-
tpu_topology=config.tpu_topology,
75+
LABELS_TO_UPDATE = (
76+
{"env": "prod"} if composer_env.is_prod_env() else {"env": "dev"}
9277
)
9378

9479
# Keyword arguments are generated dynamically at runtime (pylint does not
9580
# know this signature).
9681
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9782
group_id=f"v{config.tpu_version.value}"
9883
):
84+
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
85+
task_id="build_node_pool_info_from_gcs_yaml"
86+
)(
87+
gcs_path=GCS_CONFIG_PATH,
88+
dag_name="multi_host_nodepool_rollback",
89+
is_prod=composer_env.is_prod_env(),
90+
machine_type=config.machine_version.value,
91+
tpu_topology=config.tpu_topology,
92+
)
93+
9994
create_node_pool = node_pool.create.override(owner=test_owner.QUINN_M)(
10095
node_pool=node_pool_info,
101-
reservation="cloudtpu-20251107233000-1246578561",
10296
)
10397

10498
wait_node_pool_available = node_pool.wait_for_availability(
@@ -127,7 +121,8 @@
127121
# Airflow uses >> for task chaining, which is pointless for pylint.
128122
# pylint: disable=pointless-statement
129123
(
130-
create_node_pool
124+
node_pool_info
125+
>> create_node_pool
131126
>> wait_node_pool_available
132127
>> rollback_node_pool
133128
>> wait_node_pool_unavailable

dags/tpu_observability/node_pool_status.py

Lines changed: 36 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -14,19 +14,18 @@
1414

1515
"""A DAG to validate the status of a GKE node pool through its lifecycle."""
1616

17-
import copy
1817
import datetime
1918

2019
from airflow import models
21-
from airflow.utils.trigger_rule import TriggerRule
20+
from airflow.decorators import task
2221
from airflow.utils.task_group import TaskGroup
22+
from airflow.utils.trigger_rule import TriggerRule
2323

2424
from dags import composer_env
2525
from dags.common import test_owner
2626
from dags.map_reproducibility.utils import constants
27-
from dags.common.vm_resource import Region, Zone
27+
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
2828
from dags.tpu_observability.utils import node_pool_util as node_pool
29-
from dags.tpu_observability.configs.common import MachineConfigMap
3029

3130

3231
# Keyword arguments are generated dynamically at runtime (pylint does not
@@ -62,46 +61,48 @@
6261
) as dag:
6362
for machine in MachineConfigMap:
6463
config = machine.value
65-
cluster_name = "tpu-observability-automation"
66-
cluster_name += "-prod" if composer_env.is_prod_env() else "-dev"
67-
node_pool_info = node_pool.Info(
68-
project_id=models.Variable.get("PROJECT_ID", default_var="cienet-cmcs"),
69-
cluster_name=cluster_name,
70-
node_pool_name=models.Variable.get(
71-
"NODE_POOL_NAME", default_var="node-pool-status-v6e-autotest"
72-
),
73-
location=models.Variable.get(
74-
"LOCATION", default_var=Region.US_CENTRAL1.value
75-
),
76-
node_locations=models.Variable.get(
77-
"NODE_LOCATIONS", default_var=Zone.US_CENTRAL1_B.value
78-
),
79-
num_nodes=models.Variable.get("NUM_NODES", default_var=4),
80-
machine_type=config.machine_version.value,
81-
tpu_topology=config.tpu_topology,
82-
)
83-
84-
problematic_node_pool_info = copy.deepcopy(node_pool_info)
85-
problematic_node_pool_info.node_pool_name += "-wrong"
86-
# Choosing a region that is different from the cluster location but still
87-
# compatible with the specified TPU cause the cluster creation to fail
88-
# due to mismatched node locations.
89-
problematic_node_pool_info.node_locations = models.Variable.get(
90-
"WRONG_NODE_LOCATION", default_var=Zone.ASIA_EAST1_C.value
91-
)
64+
65+
@task
66+
def generate_problematic_node_pool_name(
67+
node_pool_info: node_pool.Info,
68+
) -> str:
69+
"""Generates a problematic node pool name."""
70+
return f"{node_pool_info.node_pool_name}-x"
71+
72+
@task
73+
def generate_problematic_node_location(
74+
node_pool_info: node_pool.Info,
75+
) -> str:
76+
"""Generates a problematic node location."""
77+
return f"{node_pool_info.location}-c"
9278

9379
# Keyword arguments are generated dynamically at runtime (pylint does not
9480
# know this signature).
9581
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9682
group_id=f"v{config.tpu_version.value}"
9783
):
84+
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
85+
task_id="build_node_pool_info_from_gcs_yaml"
86+
)(
87+
gcs_path=GCS_CONFIG_PATH,
88+
dag_name="gke_node_pool_status",
89+
is_prod=composer_env.is_prod_env(),
90+
machine_type=config.machine_version.value,
91+
tpu_topology=config.tpu_topology,
92+
)
93+
94+
problematic_node_pool_info = node_pool.copy_node_pool_info_with_override(
95+
info=node_pool_info,
96+
node_pool_name=generate_problematic_node_pool_name(node_pool_info),
97+
node_locations=generate_problematic_node_location(node_pool_info),
98+
)
99+
98100
task_id = "create_node_pool"
99101
create_node_pool = node_pool.create.override(
100102
task_id=task_id,
101103
owner=test_owner.YUNA_T,
102104
)(
103105
node_pool=node_pool_info,
104-
reservation="cloudtpu-20251107233000-1246578561",
105106
)
106107

107108
task_id = "wait_for_provisioning"
@@ -175,7 +176,9 @@
175176
# Airflow uses >> for task chaining, which is pointless for pylint.
176177
# pylint: disable=pointless-statement
177178
normal_flow = (
178-
create_node_pool
179+
node_pool_info
180+
>> problematic_node_pool_info
181+
>> create_node_pool
179182
>> wait_for_provisioning
180183
>> wait_for_running
181184
>> delete_node

dags/tpu_observability/node_pool_ttr_update_label.py

Lines changed: 13 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222

2323
from dags import composer_env
2424
from dags.common.vm_resource import Region, Zone
25-
from dags.tpu_observability.configs.common import MachineConfigMap
25+
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
2626
from dags.tpu_observability.utils import node_pool_util as node_pool
2727

2828

@@ -64,33 +64,22 @@
6464
) as dag:
6565
for machine in MachineConfigMap:
6666
config = machine.value
67-
cluster_name = "tpu-observability-automation"
68-
cluster_name += "-prod" if composer_env.is_prod_env() else "-dev"
69-
node_pool_info = node_pool.Info(
70-
project_id=models.Variable.get("PROJECT_ID", default_var="cienet-cmcs"),
71-
cluster_name=cluster_name,
72-
node_pool_name=models.Variable.get(
73-
"NODE_POOL_NAME",
74-
default_var="ttr-update-label-v6e-autotest",
75-
),
76-
location=models.Variable.get(
77-
"LOCATION", default_var=Region.US_CENTRAL1.value
78-
),
79-
node_locations=models.Variable.get(
80-
"NODE_LOCATIONS", default_var=Zone.US_CENTRAL1_B.value
81-
),
82-
num_nodes=models.Variable.get("NUM_NODES", default_var=4),
83-
machine_type=config.machine_version.value,
84-
tpu_topology=config.tpu_topology,
85-
)
86-
8767
LABELS_TO_UPDATE = {"test_key": "test_val"}
8868

8969
with TaskGroup(group_id=f"v{config.tpu_version.value}"):
70+
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
71+
task_id="build_node_pool_info_from_gcs_yaml"
72+
)(
73+
gcs_path=GCS_CONFIG_PATH,
74+
dag_name="node_pool_ttr_update_label",
75+
is_prod=composer_env.is_prod_env(),
76+
machine_type=config.machine_version.value,
77+
tpu_topology=config.tpu_topology,
78+
)
79+
9080
task_id = "create_node_pool"
9181
create_node_pool = node_pool.create.override(task_id=task_id)(
9282
node_pool=node_pool_info,
93-
reservation="cloudtpu-20251107233000-1246578561",
9483
)
9584

9685
task_id = "wait_for_provisioning"
@@ -126,7 +115,8 @@
126115
)
127116

128117
_ = (
129-
create_node_pool
118+
node_pool_info
119+
>> create_node_pool
130120
>> wait_for_provisioning
131121
>> wait_for_running
132122
>> update_node_pool_label

0 commit comments

Comments
 (0)