Skip to content

Commit 3bd44a5

Browse files
committed
refactor: encapsulate JobSet metric mapping into JobSetHealthiness enum
1 parent c52b1c5 commit 3bd44a5

2 files changed

Lines changed: 58 additions & 53 deletions

File tree

dags/tpu_observability/jobset_healthiness_validation.py

Lines changed: 29 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -88,42 +88,40 @@
8888
with TaskGroup( # pylint: disable=unexpected-keyword-arg
8989
group_id=f"v{config.tpu_version.value}"
9090
):
91-
selector = jobset.generate_node_pool_selector(
92-
"jobset-healthiness-validation"
91+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
92+
gcs_path=GCS_CONFIG_PATH,
93+
dag_name=DAG_ID,
94+
is_prod=composer_env.is_prod_env(),
95+
machine_type=config.machine_version.value,
96+
tpu_topology=config.tpu_topology,
9397
)
9498

9599
jobset_config = jobset.build_jobset_from_gcs_yaml(
96100
gcs_path=GCS_JOBSET_CONFIG_PATH,
97101
dag_name=DAG_ID,
98-
node_pool_selector=selector,
99102
)
100103

101-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
102-
task_id="build_node_pool_info_from_gcs_yaml"
103-
)(
104-
gcs_path=GCS_CONFIG_PATH,
105-
dag_name=DAG_ID,
106-
is_prod=composer_env.is_prod_env(),
107-
machine_type=config.machine_version.value,
108-
tpu_topology=config.tpu_topology,
109-
node_pool_selector=selector,
110-
)
104+
selector = jobset.generate_node_pool_selector(DAG_ID)
105+
jobset_name = jobset.generate_jobset_name(jobset_config.dag_id_prefix)
111106

112107
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
113108
node_pool=cluster_info,
109+
node_pool_selector=selector,
114110
)
115111

116112
startup = jobset.create_jobset_startup_tasks(
117113
node_pool=cluster_info,
118114
jobset_config=jobset_config,
115+
jobset_name=jobset_name,
116+
node_pool_selector=selector,
119117
workload_type=Workload.JAX_TPU_BENCHMARK,
120118
)
121119

122120
with TaskGroup(group_id="validate_running_metrics") as validate_running:
123121
running_metrics = [
124-
(JobSetHealthiness.SPECIFIED, "USE_CONFIG_REPLICAS"),
125-
(JobSetHealthiness.ACTIVE, "USE_CONFIG_REPLICAS"),
126-
(JobSetHealthiness.READY, "USE_CONFIG_REPLICAS"),
122+
(JobSetHealthiness.SPECIFIED, jobset_config.replicas),
123+
(JobSetHealthiness.ACTIVE, jobset_config.replicas),
124+
(JobSetHealthiness.READY, jobset_config.replicas),
127125
(JobSetHealthiness.FAILED, 0),
128126
(JobSetHealthiness.SUCCEEDED, 0),
129127
(JobSetHealthiness.SUSPENDED, 0),
@@ -135,22 +133,23 @@
135133
metric_name=status,
136134
expected_value=expected,
137135
node_pool=cluster_info,
138-
jobset_config=jobset_config,
136+
jobset_name=jobset_name,
139137
)
140138

141139
suspend_action = jobset.suspended_jobset.override(
142140
task_id="suspend_jobset"
143141
)(
144142
node_pool=cluster_info,
145143
jobset_config=jobset_config,
144+
jobset_name=jobset_name,
146145
)
147146

148147
with TaskGroup(
149148
group_id="validate_suspended_metrics"
150149
) as validate_suspended:
151150
suspended_metrics = [
152151
(JobSetHealthiness.ACTIVE, 0),
153-
(JobSetHealthiness.SUSPENDED, "USE_CONFIG_REPLICAS"),
152+
(JobSetHealthiness.SUSPENDED, jobset_config.replicas),
154153
]
155154
for status, expected in suspended_metrics:
156155
jobset.wait_for_jobset_metrics.override(
@@ -159,12 +158,13 @@
159158
metric_name=status,
160159
expected_value=expected,
161160
node_pool=cluster_info,
162-
jobset_config=jobset_config,
161+
jobset_name=jobset_name,
163162
)
164163

165164
resume_action = jobset.resume_jobset.override(task_id="resume_jobset")(
166165
node_pool=cluster_info,
167166
jobset_config=jobset_config,
167+
jobset_name=jobset_name,
168168
)
169169

170170
with TaskGroup(group_id="inject_and_validate_success") as success_test:
@@ -173,23 +173,25 @@
173173
)(
174174
node_pool=cluster_info,
175175
jobset_config=jobset_config,
176+
jobset_name=jobset_name,
176177
)
177178

178179
start_success_job = jobset.run_workload.override(
179180
task_id="start_success_job"
180181
)(
181182
node_pool=cluster_info,
182183
jobset_config=jobset_config,
184+
jobset_name=jobset_name,
183185
workload_type=SUCCESS_WORKLOAD,
184186
)
185187

186188
validate_succeeded_metric = jobset.wait_for_jobset_metrics.override(
187189
task_id="wait_for_succeeded_count"
188190
)(
189191
metric_name=JobSetHealthiness.SUCCEEDED,
190-
expected_value="USE_CONFIG_REPLICAS",
192+
expected_value=jobset_config.replicas,
191193
node_pool=cluster_info,
192-
jobset_config=jobset_config,
194+
jobset_name=jobset_name,
193195
)
194196

195197
chain(cleanup_for_success, start_success_job, validate_succeeded_metric)
@@ -200,21 +202,23 @@
200202
)(
201203
node_pool=cluster_info,
202204
jobset_config=jobset_config,
205+
jobset_name=jobset_name,
203206
)
204207

205208
start_fail_job = jobset.run_workload.override(task_id="start_fail_job")(
206209
node_pool=cluster_info,
207210
jobset_config=jobset_config,
211+
jobset_name=jobset_name,
208212
workload_type=FAIL_WORKLOAD,
209213
)
210214

211215
validate_failed_metric = jobset.wait_for_jobset_metrics.override(
212216
task_id="wait_for_failed_count"
213217
)(
214218
metric_name=JobSetHealthiness.FAILED,
215-
expected_value="USE_CONFIG_REPLICAS",
219+
expected_value=jobset_config.replicas,
216220
node_pool=cluster_info,
217-
jobset_config=jobset_config,
221+
jobset_name=jobset_name,
218222
)
219223

220224
chain(cleanup_for_failure, start_fail_job, validate_failed_metric)
@@ -224,6 +228,7 @@
224228
)(
225229
node_pool=cluster_info,
226230
jobset_config=jobset_config,
231+
jobset_name=jobset_name,
227232
).as_teardown(
228233
setups=startup.jobset_start_time
229234
)
@@ -236,8 +241,7 @@
236241

237242
chain(
238243
selector,
239-
jobset_config,
240-
cluster_info,
244+
jobset_name,
241245
create_node_pool,
242246
*startup.tasks,
243247
validate_running,

dags/tpu_observability/utils/jobset_util.py

Lines changed: 29 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -345,6 +345,18 @@ class JobSetHealthiness(enum.Enum):
345345
SUCCEEDED = "succeeded"
346346
SPECIFIED = "specified"
347347

348+
@property
349+
def healthiness_metric_type(self) -> str:
350+
mapping = {
351+
"ready": "prometheus.googleapis.com/kube_jobset_ready_replicas/gauge",
352+
"active": "prometheus.googleapis.com/kube_jobset_active_replicas/gauge",
353+
"suspended": "prometheus.googleapis.com/kube_jobset_suspended_replicas/gauge",
354+
"succeeded": "prometheus.googleapis.com/kube_jobset_succeeded_replicas/gauge",
355+
"failed": "prometheus.googleapis.com/kube_jobset_failed_replicas/gauge",
356+
"specified": "prometheus.googleapis.com/kube_jobset_specified_replicas/gauge",
357+
}
358+
return mapping.get(self.value)
359+
348360

349361
class Command:
350362
"""
@@ -868,7 +880,9 @@ def operate_pod(
868880

869881

870882
@task
871-
def suspended_jobset(node_pool: node_pool_info, jobset_config: JobSet):
883+
def suspended_jobset(
884+
node_pool: node_pool_info, jobset_config: JobSet, jobset_name: str
885+
):
872886
"""
873887
Suspend a jobset from the GKE cluster.
874888
@@ -889,7 +903,7 @@ def suspended_jobset(node_pool: node_pool_info, jobset_config: JobSet):
889903
Command.get_credentials_command(node_pool),
890904
Command.k8s_suspend_jobset_command(
891905
temp_config_file.name,
892-
jobset_config.jobset_name,
906+
jobset_name,
893907
jobset_config.namespace,
894908
),
895909
])
@@ -898,7 +912,9 @@ def suspended_jobset(node_pool: node_pool_info, jobset_config: JobSet):
898912

899913

900914
@task
901-
def resume_jobset(node_pool: node_pool_info, jobset_config: JobSet):
915+
def resume_jobset(
916+
node_pool: node_pool_info, jobset_config: JobSet, jobset_name: str
917+
):
902918
"""
903919
Resume a jobset from the GKE cluster.
904920
@@ -919,7 +935,7 @@ def resume_jobset(node_pool: node_pool_info, jobset_config: JobSet):
919935
Command.get_credentials_command(node_pool),
920936
Command.k8s_resume_jobset_command(
921937
temp_config_file.name,
922-
jobset_config.jobset_name,
938+
jobset_name,
923939
jobset_config.namespace,
924940
),
925941
])
@@ -994,9 +1010,9 @@ def wait_for_jobset_started(
9941010
@task.sensor(poke_interval=60, timeout=3600, mode="poke")
9951011
def wait_for_jobset_metrics(
9961012
metric_name: JobSetHealthiness,
997-
expected_value: any,
1013+
expected_value: int,
9981014
node_pool: node_pool_info,
999-
jobset_config: JobSet,
1015+
jobset_name: str,
10001016
start_time: TimeUtil = None,
10011017
) -> bool:
10021018
"""Polls Cloud Monitoring for a specific JobSet replicated job metric.
@@ -1023,25 +1039,10 @@ def wait_for_jobset_metrics(
10231039
bool: True if the current metric value matches the expected value,
10241040
False otherwise.
10251041
"""
1026-
final_expected = expected_value
1027-
if expected_value == "USE_CONFIG_REPLICAS":
1028-
final_expected = jobset_config.replicas
1029-
1030-
metric_mapping = {
1031-
"ready": "prometheus.googleapis.com/kube_jobset_ready_replicas/gauge",
1032-
"active": "prometheus.googleapis.com/kube_jobset_active_replicas/gauge",
1033-
"suspended": "prometheus.googleapis.com/kube_jobset_suspended_replicas/gauge",
1034-
"succeeded": "prometheus.googleapis.com/kube_jobset_succeeded_replicas/gauge",
1035-
"failed": "prometheus.googleapis.com/kube_jobset_failed_replicas/gauge",
1036-
"specified": "prometheus.googleapis.com/kube_jobset_specified_replicas/gauge",
1037-
}
10381042

1039-
name_str = (
1040-
metric_name.value
1041-
if hasattr(metric_name, "value")
1042-
else str(metric_name).lower()
1043-
)
1044-
metric_type = metric_mapping.get(name_str)
1043+
metric_type = metric_name.healthiness_metric_type
1044+
name_str = metric_name.value
1045+
10451046
query_start = (
10461047
start_time if start_time else TimeUtil.now() - timedelta(minutes=60)
10471048
)
@@ -1052,7 +1053,7 @@ def wait_for_jobset_metrics(
10521053
f'metric.type="{metric_type}" '
10531054
f'resource.type="prometheus_target" '
10541055
f'resource.labels.cluster="{node_pool.cluster_name}" '
1055-
f'metric.labels.jobset_name="{jobset_config.jobset_name}"'
1056+
f'metric.labels.jobset_name="{jobset_name}"'
10561057
),
10571058
start_time=query_start,
10581059
end_time=TimeUtil.now(),
@@ -1071,11 +1072,11 @@ def wait_for_jobset_metrics(
10711072
latest_value = float(point_value.int64_value)
10721073

10731074
logging.info(
1074-
f"Metric {name_str} for JobSet {jobset_config.jobset_name}: "
1075-
f"current={latest_value}, expected={final_expected}"
1075+
f"Metric {name_str} for JobSet {jobset_name}: "
1076+
f"current={latest_value}, expected={expected_value}"
10761077
)
10771078

1078-
return float(latest_value) == float(final_expected)
1079+
return float(latest_value) == float(expected_value)
10791080

10801081

10811082
@task.sensor(poke_interval=60, timeout=3600, mode="poke")

0 commit comments

Comments
 (0)