2121from airflow .utils .trigger_rule import TriggerRule
2222from airflow .utils .task_group import TaskGroup
2323
24-
2524from airflow .decorators import task
2625
2726from dags import composer_env
@@ -76,7 +75,7 @@ def check_nodes_number(
7675# Keyword arguments are generated dynamically at runtime (pylint does not
7776# know this signature).
7877with models .DAG ( # pylint: disable=unexpected-keyword-arg
79- dag_id = "jobset_ttr_drain_restart" ,
78+ dag_id = DAG_ID ,
8079 start_date = datetime .datetime (2025 , 8 , 10 ),
8180 schedule = SCHEDULE if composer_env .is_prod_env () else None ,
8281 dagrun_timeout = DAGRUN_TIMEOUT ,
@@ -125,38 +124,33 @@ def check_nodes_number(
125124 with TaskGroup ( # pylint: disable=unexpected-keyword-arg
126125 group_id = f"v{ config .tpu_version .value } "
127126 ):
127+ selector = jobset .generate_node_pool_selector (DAG_ID )
128+
128129 jobset_config = jobset .build_jobset_from_gcs_yaml (
129130 gcs_path = GCS_JOBSET_CONFIG_PATH ,
130- dag_name = "jobset_ttr_drain_restart" ,
131+ dag_name = DAG_ID ,
132+ node_pool_selector = selector ,
131133 )
132134
133- cluster_info = node_pool .build_node_pool_info_from_gcs_yaml .override (
134- task_id = "build_node_pool_info_from_gcs_yaml"
135- )(
135+ cluster_info = node_pool .build_node_pool_info_from_gcs_yaml (
136136 gcs_path = GCS_CONFIG_PATH ,
137- dag_name = "jobset_ttr_drain_restart" ,
137+ dag_name = DAG_ID ,
138138 is_prod = composer_env .is_prod_env (),
139139 machine_type = config .machine_version .value ,
140140 tpu_topology = config .tpu_topology ,
141+ node_pool_selector = selector ,
141142 )
142143
143144 create_node_pool = node_pool .create .override (task_id = "create_node_pool" )(
144145 node_pool = cluster_info ,
145146 )
146147
147- start_workload = jobset .run_workload . override ( task_id = "start_workload" ) (
148+ startup = jobset .create_jobset_startup_tasks (
148149 node_pool = cluster_info ,
149150 jobset_config = jobset_config ,
150151 workload_type = Workload .JAX_TPU_BENCHMARK ,
151152 )
152153
153- ensure_all_pods_running = jobset .wait_for_all_pods_running .override (
154- task_id = "ensure_all_pods_running"
155- )(
156- node_pool = cluster_info ,
157- jobset_config = jobset_config ,
158- )
159-
160154 select_node = node_pool .draw_random_node .override (task_id = "select_node" )(
161155 node_pool = cluster_info
162156 )
@@ -193,7 +187,7 @@ def check_nodes_number(
193187 node_pool = cluster_info ,
194188 jobset_config = jobset_config ,
195189 ).as_teardown (
196- setups = start_workload
190+ setups = startup . jobset_start_time
197191 )
198192
199193 cleanup_node_pool = node_pool .delete .override (
@@ -203,11 +197,9 @@ def check_nodes_number(
203197 )
204198
205199 chain (
206- jobset_config ,
207- cluster_info ,
200+ selector ,
208201 create_node_pool ,
209- start_workload ,
210- ensure_all_pods_running ,
202+ * startup .tasks ,
211203 select_node ,
212204 drained_node ,
213205 check_nodes_number ,
0 commit comments