Skip to content

Commit b1a9112

Browse files
authored
Fix JobSet Pod Discovery and Implement GCS-Driven JobSet Configuration (GoogleCloudPlatform#1170)
This change addresses reliability issues in JobSet management by introducing a configuration-driven lifecycle and transitioning to official Kubernetes label selectors for pod discovery.
1 parent e5a0b4c commit b1a9112

7 files changed

Lines changed: 244 additions & 201 deletions

File tree

dags/tpu_observability/configs/common.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,3 +24,7 @@ class MachineConfigMap(enum.Enum):
2424
GCS_CONFIG_PATH = (
2525
"gs://ml-auto-solutions-dag-configs/tpu_observability/dag_config.yaml"
2626
)
27+
28+
GCS_JOBSET_CONFIG_PATH = (
29+
"gs://ml-auto-solutions-dag-configs/tpu_observability/jobset_config.yaml"
30+
)

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 19 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,12 @@
2525
from dags.tpu_observability.utils import jobset_util as jobset
2626
from dags.tpu_observability.utils import node_pool_util as node_pool
2727
from dags.tpu_observability.utils.jobset_util import JobSet, Workload
28-
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
28+
from dags.tpu_observability.configs.common import (
29+
MachineConfigMap,
30+
GCS_CONFIG_PATH,
31+
GCS_JOBSET_CONFIG_PATH,
32+
)
33+
2934

3035
_DISK_SIZE_INCREMENT = 100
3136

@@ -47,7 +52,8 @@
4752
],
4853
description=(
4954
"This DAG tests the JobSet time-to-recover metric by triggering a "
50-
"node pool disk resize, then polls the metric to check if it is updated."
55+
"node pool disk resize, then polls the metric to check "
56+
"if it is updated."
5157
),
5258
doc_md="""
5359
# JobSet Time-To-Recover (TTR) Test Using Node Pool Disk Resize
@@ -73,27 +79,16 @@
7379
for machine in MachineConfigMap:
7480
config = machine.value
7581

76-
jobset_config = JobSet(
77-
jobset_name="ttr-res-v6e",
78-
namespace="default",
79-
max_restarts=10,
80-
replicated_job_name="tpu-job-slice",
81-
replicas=1,
82-
backoff_limit=0,
83-
completions=4,
84-
parallelism=4,
85-
tpu_accelerator_type="tpu-v6e-slice",
86-
tpu_topology="4x4",
87-
container_name="jax-tpu-worker",
88-
image="python:3.11",
89-
tpu_cores_per_pod=4,
90-
)
91-
9282
# Keyword arguments are generated dynamically at runtime (pylint does not
9383
# know this signature).
9484
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9585
group_id=f"v{config.tpu_version.value}"
9686
):
87+
jobset_config = jobset.build_jobset_from_gcs_yaml(
88+
gcs_path=GCS_JOBSET_CONFIG_PATH,
89+
dag_name="jobset_ttr_node_pool_resize",
90+
)
91+
9792
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
9893
task_id="build_node_pool_info_from_gcs_yaml"
9994
)(
@@ -110,17 +105,15 @@
110105

111106
start_workload = jobset.run_workload.override(task_id="start_workload")(
112107
node_pool=cluster_info,
113-
yaml_config=jobset_config.generate_yaml(
114-
workload_script=Workload.JAX_TPU_BENCHMARK
115-
),
116-
namespace=jobset_config.namespace,
108+
jobset_config=jobset_config,
109+
workload_type=Workload.JAX_TPU_BENCHMARK,
117110
)
118111

119112
ensure_all_pods_running = jobset.wait_for_all_pods_running.override(
120113
task_id="ensure_all_pods_running"
121114
)(
122-
num_pods=(jobset_config.replicas * jobset_config.parallelism),
123115
node_pool=cluster_info,
116+
jobset_config=jobset_config,
124117
)
125118

126119
node_pool_resize = node_pool.update.override(task_id="node_pool_resize")(
@@ -134,16 +127,14 @@
134127
task_id="wait_for_jobset_ttr_to_be_found"
135128
)(
136129
node_pool=cluster_info,
137-
jobset_name=jobset_config.jobset_name,
138-
start_time=node_pool_resize,
130+
jobset_config=jobset_config,
139131
)
140132

141133
cleanup_workload = jobset.end_workload.override(
142134
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
143135
)(
144136
node_pool=cluster_info,
145-
jobset_name=jobset_config.jobset_name,
146-
namespace=jobset_config.namespace,
137+
jobset_config=jobset_config,
147138
).as_teardown(
148139
setups=start_workload
149140
)
@@ -155,6 +146,7 @@
155146
)
156147

157148
chain(
149+
jobset_config,
158150
cluster_info,
159151
create_node_pool,
160152
start_workload,

dags/tpu_observability/jobset_ttr_pod_delete.py

Lines changed: 21 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,12 @@
2424
from dags import composer_env
2525
from dags.tpu_observability.utils import jobset_util as jobset
2626
from dags.tpu_observability.utils import node_pool_util as node_pool
27-
from dags.tpu_observability.utils.jobset_util import JobSet, Workload
28-
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
27+
from dags.tpu_observability.utils.jobset_util import Workload
28+
from dags.tpu_observability.configs.common import (
29+
MachineConfigMap,
30+
GCS_CONFIG_PATH,
31+
GCS_JOBSET_CONFIG_PATH,
32+
)
2933

3034
# Keyword arguments are generated dynamically at runtime (pylint does not
3135
# know this signature).
@@ -70,27 +74,15 @@
7074
for machine in MachineConfigMap:
7175
config = machine.value
7276

73-
jobset_config = JobSet(
74-
jobset_name="ttr-delete-v6e-workload",
75-
namespace="default",
76-
max_restarts=5,
77-
replicated_job_name="tpu-job-slice",
78-
replicas=1,
79-
backoff_limit=0,
80-
completions=4,
81-
parallelism=4,
82-
tpu_accelerator_type="tpu-v6e-slice",
83-
tpu_topology="4x4",
84-
container_name="jax-tpu-worker",
85-
image="python:3.11",
86-
tpu_cores_per_pod=4,
87-
)
88-
8977
# Keyword arguments are generated dynamically at runtime (pylint does not
9078
# know this signature).
9179
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9280
group_id=f"v{config.tpu_version.value}"
9381
):
82+
jobset_config = jobset.build_jobset_from_gcs_yaml(
83+
gcs_path=GCS_JOBSET_CONFIG_PATH, dag_name="jobset_ttr_pod_delete"
84+
)
85+
9486
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
9587
task_id="build_node_pool_info_from_gcs_yaml"
9688
)(
@@ -107,37 +99,34 @@
10799

108100
start_workload = jobset.run_workload.override(task_id="start_workload")(
109101
node_pool=cluster_info,
110-
yaml_config=jobset_config.generate_yaml(
111-
workload_script=Workload.JAX_TPU_BENCHMARK
112-
),
113-
namespace=jobset_config.namespace,
102+
jobset_config=jobset_config,
103+
workload_type=Workload.JAX_TPU_BENCHMARK,
114104
)
115105

116106
ensure_all_pods_running = jobset.wait_for_all_pods_running.override(
117107
task_id="ensure_all_pods_running"
118108
)(
119-
num_pods=(jobset_config.replicas * jobset_config.parallelism),
120109
node_pool=cluster_info,
110+
jobset_config=jobset_config,
121111
)
122112

123113
delete_random_pod = jobset.delete_one_random_pod.override(
124114
task_id="delete_random_pod"
125-
)(node_pool=cluster_info, namespace=jobset_config.namespace)
115+
)(
116+
node_pool=cluster_info,
117+
jobset_config=jobset_config,
118+
)
126119

127120
wait_for_metric_upload = jobset.wait_for_jobset_ttr_to_be_found.override(
128121
task_id="wait_for_jobset_ttr_to_be_found"
129122
)(
130123
node_pool=cluster_info,
131-
jobset_name=jobset_config.jobset_name,
124+
jobset_config=jobset_config,
132125
)
133126

134127
cleanup_workload = jobset.end_workload.override(
135128
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
136-
)(
137-
node_pool=cluster_info,
138-
jobset_name=jobset_config.jobset_name,
139-
namespace=jobset_config.namespace,
140-
).as_teardown(
129+
)(node_pool=cluster_info, jobset_config=jobset_config).as_teardown(
141130
setups=start_workload
142131
)
143132

@@ -148,6 +137,8 @@
148137
)
149138

150139
chain(
140+
jobset_config,
141+
cluster_info,
151142
create_node_pool,
152143
start_workload,
153144
ensure_all_pods_running,

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 27 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,12 @@
2424
from dags import composer_env
2525
from dags.tpu_observability.utils import jobset_util as jobset
2626
from dags.tpu_observability.utils import node_pool_util as node_pool
27-
from dags.tpu_observability.utils.jobset_util import JobSet, Workload
28-
from dags.tpu_observability.configs.common import MachineConfigMap, GCS_CONFIG_PATH
27+
from dags.tpu_observability.utils.jobset_util import Workload
28+
from dags.tpu_observability.configs.common import (
29+
MachineConfigMap,
30+
GCS_CONFIG_PATH,
31+
GCS_JOBSET_CONFIG_PATH,
32+
)
2933

3034
# Keyword arguments are generated dynamically at runtime (pylint does not
3135
# know this signature).
@@ -72,27 +76,15 @@
7276
for machine in MachineConfigMap:
7377
config = machine.value
7478

75-
jobset_config = JobSet(
76-
jobset_name="ttr-rollback-v6e-workload",
77-
namespace="default",
78-
max_restarts=5,
79-
replicated_job_name="tpu-job-slice",
80-
replicas=1,
81-
backoff_limit=0,
82-
completions=4,
83-
parallelism=4,
84-
tpu_accelerator_type="tpu-v6e-slice",
85-
tpu_topology="4x4",
86-
container_name="jax-tpu-worker",
87-
image="python:3.11",
88-
tpu_cores_per_pod=4,
89-
)
90-
9179
# Keyword arguments are generated dynamically at runtime (pylint does not
9280
# know this signature).
9381
with TaskGroup( # pylint: disable=unexpected-keyword-arg
9482
group_id=f"v{config.tpu_version.value}"
9583
):
84+
jobset_config = jobset.build_jobset_from_gcs_yaml(
85+
gcs_path=GCS_JOBSET_CONFIG_PATH, dag_name="jobset_rollback_ttr"
86+
)
87+
9688
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
9789
task_id="build_node_pool_info_from_gcs_yaml"
9890
)(
@@ -103,36 +95,39 @@
10395
tpu_topology=config.tpu_topology,
10496
)
10597

106-
create_node_pool = node_pool.create(
98+
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
10799
node_pool=cluster_info,
108100
)
109101

110-
start_workload = jobset.run_workload(
102+
start_workload = jobset.run_workload.override(task_id="start_workload")(
111103
node_pool=cluster_info,
112-
yaml_config=jobset_config.generate_yaml(
113-
workload_script=Workload.JAX_TPU_BENCHMARK
114-
),
115-
namespace=jobset_config.namespace,
104+
jobset_config=jobset_config,
105+
workload_type=Workload.JAX_TPU_BENCHMARK,
116106
)
117107

118-
ensure_all_pods_running = jobset.wait_for_all_pods_running(
119-
num_pods=(jobset_config.replicas * jobset_config.parallelism),
108+
ensure_all_pods_running = jobset.wait_for_all_pods_running.override(
109+
task_id="ensure_all_pods_running"
110+
)(
120111
node_pool=cluster_info,
112+
jobset_config=jobset_config,
121113
)
122114

123-
rollback_node_pool = node_pool.rollback(node_pool=cluster_info)
115+
rollback_node_pool = node_pool.rollback.override(
116+
task_id="rollback_node_pool"
117+
)(node_pool=cluster_info)
124118

125-
wait_for_metric_upload = jobset.wait_for_jobset_ttr_to_be_found(
119+
wait_for_metric_upload = jobset.wait_for_jobset_ttr_to_be_found.override(
120+
task_id="wait_for_jobset_ttr_to_be_found"
121+
)(
126122
node_pool=cluster_info,
127-
jobset_name=jobset_config.jobset_name,
123+
jobset_config=jobset_config,
128124
)
129125

130126
cleanup_workload = jobset.end_workload.override(
131127
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
132128
)(
133129
node_pool=cluster_info,
134-
jobset_name=jobset_config.jobset_name,
135-
namespace=jobset_config.namespace,
130+
jobset_config=jobset_config,
136131
).as_teardown(
137132
setups=start_workload
138133
)
@@ -144,6 +139,7 @@
144139
)
145140

146141
chain(
142+
jobset_config,
147143
cluster_info,
148144
create_node_pool,
149145
start_workload,

0 commit comments

Comments
 (0)