Skip to content

Commit 0ff7f02

Browse files
fix: Add additional label/selector to make sure JobSet is dispatched to desired NodePool (GoogleCloudPlatform#1187)
This change fixes a critical scheduling issue where JobSet pods were being dispatched to arbitrary node pools instead of the intended ones, and extends the fix to support multi-node-pool environments. Problem: Previously, the JobSet YAML template had no nodeSelector for the GKE node pool. In environments with multiple node pools, Kubernetes would schedule pods on any available nodes. This led to cases where a Rollback was performed on Node Pool A, but the JobSet pods were actually running on Node Pool B, resulting in inaccurate TTR (Time To Recovery) metrics. Additionally, for DAGs like `tpu_info_format_validation_dag` that create two node pools for the same workload, pinning to a single node pool name would leave the second pool unused.
1 parent 08d6cff commit 0ff7f02

8 files changed

Lines changed: 80 additions & 6 deletions

dags/tpu_observability/jobset_ttr_kill_process.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -127,9 +127,12 @@ def kill_tpu_pod_workload(info: node_pool.Info, pod_name: str) -> None:
127127
with TaskGroup( # pylint: disable=unexpected-keyword-arg
128128
group_id=f"v{config.tpu_version.value}"
129129
):
130+
selector = jobset.generate_node_pool_selector("jobset-ttr-kill-process")
131+
130132
jobset_config = jobset.build_jobset_from_gcs_yaml(
131133
gcs_path=GCS_JOBSET_CONFIG_PATH,
132134
dag_name="jobset_ttr_kill_process",
135+
node_pool_selector=selector,
133136
)
134137

135138
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -140,6 +143,7 @@ def kill_tpu_pod_workload(info: node_pool.Info, pod_name: str) -> None:
140143
is_prod=composer_env.is_prod_env(),
141144
machine_type=config.machine_version.value,
142145
tpu_topology=config.tpu_topology,
146+
node_pool_selector=selector,
143147
)
144148

145149
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
@@ -190,6 +194,7 @@ def kill_tpu_pod_workload(info: node_pool.Info, pod_name: str) -> None:
190194
)
191195

192196
chain(
197+
selector,
193198
jobset_config,
194199
cluster_info,
195200
create_node_pool,

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,9 +89,14 @@
8989
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9090
group_id=f"v{config.tpu_version.value}"
9191
):
92+
selector = jobset.generate_node_pool_selector(
93+
"jobset-ttr-node-pool-resize"
94+
)
95+
9296
jobset_config = jobset.build_jobset_from_gcs_yaml(
9397
gcs_path=GCS_JOBSET_CONFIG_PATH,
9498
dag_name="jobset_ttr_node_pool_resize",
99+
node_pool_selector=selector,
95100
)
96101

97102
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -102,6 +107,7 @@
102107
is_prod=composer_env.is_prod_env(),
103108
machine_type=config.machine_version.value,
104109
tpu_topology=config.tpu_topology,
110+
node_pool_selector=selector,
105111
)
106112

107113
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
@@ -151,6 +157,7 @@
151157
)
152158

153159
chain(
160+
selector,
154161
jobset_config,
155162
cluster_info,
156163
create_node_pool,

dags/tpu_observability/jobset_ttr_pod_delete.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,8 +86,12 @@
8686
with TaskGroup( # pylint: disable=unexpected-keyword-arg
8787
group_id=f"v{config.tpu_version.value}"
8888
):
89+
selector = jobset.generate_node_pool_selector("jobset-ttr-pod-delete")
90+
8991
jobset_config = jobset.build_jobset_from_gcs_yaml(
90-
gcs_path=GCS_JOBSET_CONFIG_PATH, dag_name="jobset_ttr_pod_delete"
92+
gcs_path=GCS_JOBSET_CONFIG_PATH,
93+
dag_name="jobset_ttr_pod_delete",
94+
node_pool_selector=selector,
9195
)
9296

9397
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -98,6 +102,7 @@
98102
is_prod=composer_env.is_prod_env(),
99103
machine_type=config.machine_version.value,
100104
tpu_topology=config.tpu_topology,
105+
node_pool_selector=selector,
101106
)
102107

103108
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
@@ -144,6 +149,7 @@
144149
)
145150

146151
chain(
152+
selector,
147153
jobset_config,
148154
cluster_info,
149155
create_node_pool,

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -88,8 +88,12 @@
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("jobset-rollback-ttr")
92+
9193
jobset_config = jobset.build_jobset_from_gcs_yaml(
92-
gcs_path=GCS_JOBSET_CONFIG_PATH, dag_name="jobset_rollback_ttr"
94+
gcs_path=GCS_JOBSET_CONFIG_PATH,
95+
dag_name="jobset_rollback_ttr",
96+
node_pool_selector=selector,
9397
)
9498

9599
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -100,6 +104,7 @@
100104
is_prod=composer_env.is_prod_env(),
101105
machine_type=config.machine_version.value,
102106
tpu_topology=config.tpu_topology,
107+
node_pool_selector=selector,
103108
)
104109

105110
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
@@ -146,6 +151,7 @@
146151
)
147152

148153
chain(
154+
selector,
149155
jobset_config,
150156
cluster_info,
151157
create_node_pool,

dags/tpu_observability/tpu_info_format_validation_dags.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -346,6 +346,10 @@ def generate_second_node_pool_name(
346346
"""Generates a second node pool name."""
347347
return f"{node_pool_info.node_pool_name}-2"
348348

349+
selector = jobset.generate_node_pool_selector(
350+
"tpu-info-format-validation-dag"
351+
)
352+
349353
# Keyword arguments are generated dynamically at runtime (pylint does not
350354
# know this signature).
351355
with TaskGroup( # pylint: disable=unexpected-keyword-arg
@@ -354,6 +358,7 @@ def generate_second_node_pool_name(
354358
jobset_config = jobset.build_jobset_from_gcs_yaml(
355359
gcs_path=GCS_JOBSET_CONFIG_PATH,
356360
dag_name="tpu_info_format_validation_dag",
361+
node_pool_selector=selector,
357362
)
358363

359364
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -364,6 +369,7 @@ def generate_second_node_pool_name(
364369
is_prod=composer_env.is_prod_env(),
365370
machine_type=config.machine_version.value,
366371
tpu_topology=config.tpu_topology,
372+
node_pool_selector=selector,
367373
)
368374

369375
cluster_info_2 = node_pool.copy_node_pool_info_with_override.override(
@@ -509,6 +515,7 @@ def generate_second_node_pool_name(
509515
chain(cleanup_first_node_pool, cleanup_second_node_pool)
510516

511517
chain(
518+
selector,
512519
jobset_config,
513520
cluster_info,
514521
cluster_info_2,

dags/tpu_observability/tpu_sdk_monitoring_validation_dag.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,9 +131,14 @@ def validate_monitoring_sdk(info: node_pool.Info, pod_name: str) -> None:
131131
with TaskGroup( # pylint: disable=unexpected-keyword-arg
132132
group_id=f"v{config.tpu_version.value}"
133133
):
134+
selector = jobset.generate_node_pool_selector(
135+
"tpu-sdk-monitoring-validation"
136+
)
137+
134138
jobset_config = jobset.build_jobset_from_gcs_yaml(
135139
gcs_path=GCS_JOBSET_CONFIG_PATH,
136140
dag_name="tpu_sdk_monitoring_validation",
141+
node_pool_selector=selector,
137142
)
138143

139144
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
@@ -144,6 +149,7 @@ def validate_monitoring_sdk(info: node_pool.Info, pod_name: str) -> None:
144149
is_prod=composer_env.is_prod_env(),
145150
machine_type=config.machine_version.value,
146151
tpu_topology=config.tpu_topology,
152+
node_pool_selector=selector,
147153
)
148154

149155
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
@@ -191,6 +197,7 @@ def validate_monitoring_sdk(info: node_pool.Info, pod_name: str) -> None:
191197
)
192198

193199
chain(
200+
selector,
194201
jobset_config,
195202
cluster_info,
196203
create_node_pool,

dags/tpu_observability/utils/jobset_util.py

Lines changed: 20 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -34,13 +34,26 @@
3434
from dags.tpu_observability.utils import subprocess_util as subprocess
3535
from dags.tpu_observability.utils.gcp_util import query_time_series
3636
from dags.tpu_observability.utils.node_pool_util import Info as node_pool_info
37+
from dags.tpu_observability.utils.node_pool_util import NODE_POOL_SELECTOR_KEY
3738
from dags.tpu_observability.utils.time_util import TimeUtil
38-
from google.cloud.monitoring_v3 import types
39-
import kubernetes
4039
from xlml.apis import gcs
4140
from xlml.utils import gke
4241

4342

43+
@task
44+
def generate_node_pool_selector(prefix: str) -> str:
45+
"""Generates a unique node_pool_selector value.
46+
47+
Args:
48+
prefix: An identifier for the workload type (e.g., "resize", "rollback").
49+
50+
Returns:
51+
The selector value string (e.g., "rollback-20260212123456").
52+
"""
53+
run_id = datetime.datetime.now().strftime("%Y%m%d%H%M%S")
54+
return f"{prefix}-{run_id}"
55+
56+
4457
class Workload:
4558
"""A library of predefined workload scripts for JobSet.
4659
@@ -129,7 +142,7 @@ def matmul_ultra_heavy(x, y):
129142
# pylint: disable=line-too-long
130143
_TEMPLATE = string.Template(
131144
textwrap.dedent(
132-
"""
145+
f"""
133146
apiVersion: jobset.x-k8s.io/v1alpha2
134147
kind: JobSet
135148
metadata:
@@ -153,6 +166,7 @@ def matmul_ultra_heavy(x, y):
153166
nodeSelector:
154167
cloud.google.com/gke-tpu-accelerator: $tpu_accelerator_type
155168
cloud.google.com/gke-tpu-topology: $tpu_topology
169+
{NODE_POOL_SELECTOR_KEY}: $node_pool_selector
156170
containers:
157171
- name: $container_name
158172
image: $image
@@ -212,6 +226,7 @@ class JobSet:
212226
container_name: str
213227
image: str
214228
tpu_cores_per_pod: int
229+
node_pool_selector: str
215230

216231
def generate_yaml(self, workload_script: Workload) -> str:
217232
"""Generates the final JobSet YAML content.
@@ -226,6 +241,7 @@ def generate_yaml(self, workload_script: Workload) -> str:
226241
params = dataclasses.asdict(self)
227242
params["command"] = ["bash", "-c"]
228243
params["args"] = workload_script
244+
params["node_pool_selector"] = self.node_pool_selector or ""
229245

230246
return _TEMPLATE.substitute(params)
231247

@@ -472,6 +488,7 @@ def run_workload(
472488
Args:
473489
node_pool: Configuration object with cluster details.
474490
jobset_config: The JobSet object containing YAML configuration.
491+
workload_type: The workload script to execute.
475492
Returns:
476493
The UTC time when the workload was started.
477494
"""

dags/tpu_observability/utils/node_pool_util.py

Lines changed: 20 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,16 @@
3333
from xlml.utils import composer
3434

3535

36+
NODE_POOL_SELECTOR_KEY = "tpu-observability/workload"
37+
"""The label key for binding JobSet workloads to specific GKE node pools.
38+
39+
This key is used as a Kubernetes node label to ensure pods are scheduled
40+
on the correct node pool. It is applied to both:
41+
- GKE node pools via `--node-labels` during creation
42+
- JobSet YAML via `nodeSelector` to target the labeled nodes
43+
"""
44+
45+
3646
class Status(enum.Enum):
3747
"""Enum for GKE node pool status."""
3848

@@ -71,6 +81,7 @@ class Info:
7181
num_nodes: int = None
7282
tpu_topology: str = None
7383
reservation: str = None
84+
node_pool_selector: str = None
7485

7586

7687
@task
@@ -183,7 +194,12 @@ def create(
183194
node_pool: Info,
184195
ignore_failure: bool = False,
185196
) -> None:
186-
"""Creates a GKE node pool by the given node pool information."""
197+
"""Creates a GKE node pool by the given node pool information.
198+
199+
Args:
200+
node_pool: The node pool configuration.
201+
ignore_failure: If True, command failures are ignored.
202+
"""
187203

188204
composer.log_metadata_for_xlml_dashboard({
189205
"cluster_project": node_pool.project_id,
@@ -214,6 +230,9 @@ def create(
214230
if node_pool.reservation:
215231
command += f" --reservation-affinity=specific --reservation={node_pool.reservation}"
216232

233+
if node_pool.node_pool_selector:
234+
command += f" --node-labels={NODE_POOL_SELECTOR_KEY}={node_pool.node_pool_selector}"
235+
217236
if ignore_failure:
218237
command += "2>&1 || true "
219238

0 commit comments

Comments
 (0)