|
88 | 88 | with TaskGroup( # pylint: disable=unexpected-keyword-arg |
89 | 89 | group_id=f"v{config.tpu_version.value}" |
90 | 90 | ): |
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, |
93 | 97 | ) |
94 | 98 |
|
95 | 99 | jobset_config = jobset.build_jobset_from_gcs_yaml( |
96 | 100 | gcs_path=GCS_JOBSET_CONFIG_PATH, |
97 | 101 | dag_name=DAG_ID, |
98 | | - node_pool_selector=selector, |
99 | 102 | ) |
100 | 103 |
|
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) |
111 | 106 |
|
112 | 107 | create_node_pool = node_pool.create.override(task_id="create_node_pool")( |
113 | 108 | node_pool=cluster_info, |
| 109 | + node_pool_selector=selector, |
114 | 110 | ) |
115 | 111 |
|
116 | 112 | startup = jobset.create_jobset_startup_tasks( |
117 | 113 | node_pool=cluster_info, |
118 | 114 | jobset_config=jobset_config, |
| 115 | + jobset_name=jobset_name, |
| 116 | + node_pool_selector=selector, |
119 | 117 | workload_type=Workload.JAX_TPU_BENCHMARK, |
120 | 118 | ) |
121 | 119 |
|
122 | 120 | with TaskGroup(group_id="validate_running_metrics") as validate_running: |
123 | 121 | 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), |
127 | 125 | (JobSetHealthiness.FAILED, 0), |
128 | 126 | (JobSetHealthiness.SUCCEEDED, 0), |
129 | 127 | (JobSetHealthiness.SUSPENDED, 0), |
|
135 | 133 | metric_name=status, |
136 | 134 | expected_value=expected, |
137 | 135 | node_pool=cluster_info, |
138 | | - jobset_config=jobset_config, |
| 136 | + jobset_name=jobset_name, |
139 | 137 | ) |
140 | 138 |
|
141 | 139 | suspend_action = jobset.suspended_jobset.override( |
142 | 140 | task_id="suspend_jobset" |
143 | 141 | )( |
144 | 142 | node_pool=cluster_info, |
145 | 143 | jobset_config=jobset_config, |
| 144 | + jobset_name=jobset_name, |
146 | 145 | ) |
147 | 146 |
|
148 | 147 | with TaskGroup( |
149 | 148 | group_id="validate_suspended_metrics" |
150 | 149 | ) as validate_suspended: |
151 | 150 | suspended_metrics = [ |
152 | 151 | (JobSetHealthiness.ACTIVE, 0), |
153 | | - (JobSetHealthiness.SUSPENDED, "USE_CONFIG_REPLICAS"), |
| 152 | + (JobSetHealthiness.SUSPENDED, jobset_config.replicas), |
154 | 153 | ] |
155 | 154 | for status, expected in suspended_metrics: |
156 | 155 | jobset.wait_for_jobset_metrics.override( |
|
159 | 158 | metric_name=status, |
160 | 159 | expected_value=expected, |
161 | 160 | node_pool=cluster_info, |
162 | | - jobset_config=jobset_config, |
| 161 | + jobset_name=jobset_name, |
163 | 162 | ) |
164 | 163 |
|
165 | 164 | resume_action = jobset.resume_jobset.override(task_id="resume_jobset")( |
166 | 165 | node_pool=cluster_info, |
167 | 166 | jobset_config=jobset_config, |
| 167 | + jobset_name=jobset_name, |
168 | 168 | ) |
169 | 169 |
|
170 | 170 | with TaskGroup(group_id="inject_and_validate_success") as success_test: |
|
173 | 173 | )( |
174 | 174 | node_pool=cluster_info, |
175 | 175 | jobset_config=jobset_config, |
| 176 | + jobset_name=jobset_name, |
176 | 177 | ) |
177 | 178 |
|
178 | 179 | start_success_job = jobset.run_workload.override( |
179 | 180 | task_id="start_success_job" |
180 | 181 | )( |
181 | 182 | node_pool=cluster_info, |
182 | 183 | jobset_config=jobset_config, |
| 184 | + jobset_name=jobset_name, |
183 | 185 | workload_type=SUCCESS_WORKLOAD, |
184 | 186 | ) |
185 | 187 |
|
186 | 188 | validate_succeeded_metric = jobset.wait_for_jobset_metrics.override( |
187 | 189 | task_id="wait_for_succeeded_count" |
188 | 190 | )( |
189 | 191 | metric_name=JobSetHealthiness.SUCCEEDED, |
190 | | - expected_value="USE_CONFIG_REPLICAS", |
| 192 | + expected_value=jobset_config.replicas, |
191 | 193 | node_pool=cluster_info, |
192 | | - jobset_config=jobset_config, |
| 194 | + jobset_name=jobset_name, |
193 | 195 | ) |
194 | 196 |
|
195 | 197 | chain(cleanup_for_success, start_success_job, validate_succeeded_metric) |
|
200 | 202 | )( |
201 | 203 | node_pool=cluster_info, |
202 | 204 | jobset_config=jobset_config, |
| 205 | + jobset_name=jobset_name, |
203 | 206 | ) |
204 | 207 |
|
205 | 208 | start_fail_job = jobset.run_workload.override(task_id="start_fail_job")( |
206 | 209 | node_pool=cluster_info, |
207 | 210 | jobset_config=jobset_config, |
| 211 | + jobset_name=jobset_name, |
208 | 212 | workload_type=FAIL_WORKLOAD, |
209 | 213 | ) |
210 | 214 |
|
211 | 215 | validate_failed_metric = jobset.wait_for_jobset_metrics.override( |
212 | 216 | task_id="wait_for_failed_count" |
213 | 217 | )( |
214 | 218 | metric_name=JobSetHealthiness.FAILED, |
215 | | - expected_value="USE_CONFIG_REPLICAS", |
| 219 | + expected_value=jobset_config.replicas, |
216 | 220 | node_pool=cluster_info, |
217 | | - jobset_config=jobset_config, |
| 221 | + jobset_name=jobset_name, |
218 | 222 | ) |
219 | 223 |
|
220 | 224 | chain(cleanup_for_failure, start_fail_job, validate_failed_metric) |
|
224 | 228 | )( |
225 | 229 | node_pool=cluster_info, |
226 | 230 | jobset_config=jobset_config, |
| 231 | + jobset_name=jobset_name, |
227 | 232 | ).as_teardown( |
228 | 233 | setups=startup.jobset_start_time |
229 | 234 | ) |
|
236 | 241 |
|
237 | 242 | chain( |
238 | 243 | selector, |
239 | | - jobset_config, |
240 | | - cluster_info, |
| 244 | + jobset_name, |
241 | 245 | create_node_pool, |
242 | 246 | *startup.tasks, |
243 | 247 | validate_running, |
|
0 commit comments