Skip to content

Commit b00e4e1

Browse files
committed
Refactor to use new
-selector task -new global variable syntax -chain instead of ">>"
1 parent cb8dce9 commit b00e4e1

1 file changed

Lines changed: 35 additions & 32 deletions

File tree

dags/tpu_observability/jobset_healthiness_ready.py

Lines changed: 35 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -18,22 +18,33 @@
1818

1919
from airflow import models
2020
from airflow.decorators import task
21+
from airflow.models.baseoperator import chain
2122
from airflow.utils.trigger_rule import TriggerRule
2223
from airflow.utils.task_group import TaskGroup
2324

2425
from dags import composer_env
2526
from dags.tpu_observability.utils import jobset_util as jobset
2627
from dags.tpu_observability.utils import node_pool_util as node_pool
2728
from dags.tpu_observability.utils.jobset_util import JobSet, Workload
28-
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
29+
from dags.tpu_observability.configs.common import (
30+
MachineConfigMap,
31+
GCS_CONFIG_PATH,
32+
GCS_JOBSET_CONFIG_PATH,
33+
)
34+
from dags.common.scheduling_helper.scheduling_helper import SchedulingHelper, get_dag_timeout
2935

3036

37+
DAG_ID = "jobset_healthiness_ready"
38+
DAGRUN_TIMEOUT = get_dag_timeout(DAG_ID)
39+
SCHEDULE = SchedulingHelper.arrange_schedule_time(DAG_ID)
40+
3141
# Keyword arguments are generated dynamically at runtime (pylint does not
3242
# know this signature).
3343
with models.DAG( # pylint: disable=unexpected-keyword-arg
34-
dag_id="jobset_healthiness_ready",
44+
dag_id=DAG_ID,
3545
start_date=datetime.datetime(2025, 8, 10),
36-
schedule="30 19 * * *" if composer_env.is_prod_env() else None,
46+
schedule=SCHEDULE if composer_env.is_prod_env() else None,
47+
dagrun_timeout=DAGRUN_TIMEOUT,
3748
catchup=False,
3849
tags=[
3950
"cloud-ml-auto-solutions",
@@ -78,35 +89,28 @@ def generate_second_node_pool_name(
7889
"""Generates a second node pool name."""
7990
return f"{node_pool_info.node_pool_name}-2"
8091

81-
jobset_config = JobSet(
82-
jobset_name="jobset-healthiness-ready",
83-
namespace="default",
84-
max_restarts=0,
85-
replicated_job_name="tpu-job-slice",
86-
replicas=2,
87-
backoff_limit=0,
88-
completions=4,
89-
parallelism=4,
90-
tpu_accelerator_type="tpu-v6e-slice",
91-
tpu_topology="4x4",
92-
container_name="jax-tpu-worker",
93-
image="python:3.11",
94-
tpu_cores_per_pod=4,
95-
)
96-
9792
# Keyword arguments are generated dynamically at runtime (pylint does not
9893
# know this signature).
9994
with TaskGroup( # pylint: disable=unexpected-keyword-arg
10095
group_id=f"v{config.tpu_version.value}"
10196
):
97+
selector = jobset.generate_node_pool_selector("jobset_healthiness_ready")
98+
99+
jobset_config = jobset.build_jobset_from_gcs_yaml(
100+
gcs_path=GCS_JOBSET_CONFIG_PATH,
101+
dag_name=DAG_ID,
102+
node_pool_selector=selector,
103+
)
104+
102105
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
103106
task_id="build_node_pool_info_from_gcs_yaml"
104107
)(
105108
gcs_path=GCS_CONFIG_PATH,
106-
dag_name="jobset_healthiness_ready",
109+
dag_name=DAG_ID,
107110
is_prod=composer_env.is_prod_env(),
108111
machine_type=config.machine_version.value,
109112
tpu_topology=config.tpu_topology,
113+
node_pool_selector=selector,
110114
)
111115

112116
cluster_info_2 = node_pool.copy_node_pool_info_with_override(
@@ -182,16 +186,15 @@ def generate_second_node_pool_name(
182186
setups=create_node_pool,
183187
)
184188

185-
# Airflow uses >> for task chaining, which is pointless for pylint.
186-
# pylint: disable=pointless-statement
187-
(
188-
cluster_info
189-
>> cluster_info_2
190-
>> create_node_pool
191-
>> validate_zero_replicas
192-
>> start_workload
193-
>> validate_ready_replicas
194-
>> cleanup_workload
195-
>> cleanup_node_pool
189+
chain(
190+
selector,
191+
jobset_config,
192+
cluster_info,
193+
cluster_info_2,
194+
create_node_pool,
195+
validate_zero_replicas,
196+
start_workload,
197+
validate_ready_replicas,
198+
cleanup_workload,
199+
cleanup_node_pool,
196200
)
197-
# pylint: enable=pointless-statement

0 commit comments

Comments
 (0)