@@ -285,22 +285,26 @@ def generate_second_node_pool_name(
285285 """Generates a second node pool name."""
286286 return f"{ node_pool_info .node_pool_name } -2"
287287
288- workload_script = jobset .Workload .JAX_TPU_BENCHMARK
289-
290288 with TaskGroup (group_id = f"v{ config .tpu_version .value } " ):
289+ selector = jobset .generate_node_pool_selector (
290+ "tpu_info_metrics_verification"
291+ )
292+
291293 jobset_config = jobset .build_jobset_from_gcs_yaml (
292294 gcs_path = GCS_JOBSET_CONFIG_PATH ,
293- dag_name = "tpu_info_metrics_verification" ,
295+ dag_name = DAG_ID ,
296+ node_pool_selector = selector ,
294297 )
295298
296299 cluster_info = node_pool .build_node_pool_info_from_gcs_yaml .override (
297300 task_id = "build_node_pool_info_from_gcs_yaml"
298301 )(
299302 gcs_path = GCS_CONFIG_PATH ,
300- dag_name = "tpu_info_metrics_verification" ,
303+ dag_name = DAG_ID ,
301304 is_prod = composer_env .is_prod_env (),
302305 machine_type = config .machine_version .value ,
303306 tpu_topology = config .tpu_topology ,
307+ node_pool_selector = selector ,
304308 )
305309
306310 cluster_info_2 = node_pool .copy_node_pool_info_with_override (
@@ -417,6 +421,7 @@ def generate_second_node_pool_name(
417421 chain (cleanup_first_node_pool , cleanup_second_node_pool )
418422
419423 chain (
424+ selector ,
420425 jobset_config ,
421426 cluster_info ,
422427 cluster_info_2 ,
0 commit comments