File tree Expand file tree Collapse file tree
Expand file tree Collapse file tree Original file line number Diff line number Diff line change @@ -160,6 +160,10 @@ def check_nodes_number(
160160 node_pool_selector = selector ,
161161 )
162162
163+ with TaskGroupWithTimeout (
164+ group_id = "test" ,
165+ timeout = TEST_TIMEOUT ,
166+ ) as test :
163167 startup = jobset .create_jobset_startup_tasks (
164168 node_pool = cluster_info ,
165169 jobset_config = jobset_config ,
@@ -168,10 +172,6 @@ def check_nodes_number(
168172 workload_type = Workload .JAX_TPU_BENCHMARK ,
169173 )
170174
171- with TaskGroupWithTimeout (
172- group_id = "test" ,
173- timeout = TEST_TIMEOUT ,
174- ) as test :
175175 select_node = node_pool .draw_random_node .override (
176176 task_id = "select_node"
177177 )(node_pool = cluster_info )
@@ -207,6 +207,7 @@ def check_nodes_number(
207207 )
208208
209209 chain (
210+ * startup .tasks ,
210211 select_node ,
211212 drained_node ,
212213 check_nodes_number_task ,
Original file line number Diff line number Diff line change 125125 node_pool_selector = selector ,
126126 )
127127
128+ with TaskGroupWithTimeout (
129+ group_id = "test" ,
130+ timeout = TEST_TIMEOUT ,
131+ ) as test :
128132 startup = jobset .create_jobset_startup_tasks (
129133 node_pool = cluster_info ,
130134 jobset_config = jobset_config ,
133137 workload_type = Workload .JAX_TPU_BENCHMARK ,
134138 )
135139
136- with TaskGroupWithTimeout (
137- group_id = "test" ,
138- timeout = TEST_TIMEOUT ,
139- ) as test :
140140 node_pool_resize_start_time = node_pool .update .override (
141141 task_id = "node_pool_resize"
142142 )(
172172 )
173173
174174 chain (
175+ * startup .tasks ,
175176 node_pool_resize_start_time ,
176177 wait_for_recovery ,
177178 verify_duration ,
Original file line number Diff line number Diff line change 128128 node_pool_selector = selector ,
129129 )
130130
131+ with TaskGroupWithTimeout (
132+ group_id = "test" ,
133+ timeout = TEST_TIMEOUT ,
134+ ) as test :
131135 startup = jobset .create_jobset_startup_tasks (
132136 node_pool = cluster_info ,
133137 jobset_config = jobset_config ,
136140 workload_type = Workload .JAX_TPU_BENCHMARK ,
137141 )
138142
139- with TaskGroupWithTimeout (
140- group_id = "test" ,
141- timeout = TEST_TIMEOUT ,
142- ) as test :
143143 deletion_start_time = jobset .delete_one_random_pod .override (
144144 task_id = "delete_random_pod"
145145 )(
174174 )
175175
176176 chain (
177+ * startup .tasks ,
177178 deletion_start_time ,
178179 wait_for_recovery ,
179180 verify_duration ,
Original file line number Diff line number Diff line change 129129 node_pool_selector = selector ,
130130 )
131131
132+ with TaskGroupWithTimeout (
133+ group_id = "test" ,
134+ timeout = TEST_TIMEOUT ,
135+ ) as test :
132136 startup = jobset .create_jobset_startup_tasks (
133137 node_pool = cluster_info ,
134138 jobset_config = jobset_config ,
137141 workload_type = Workload .JAX_TPU_BENCHMARK ,
138142 )
139143
140- with TaskGroupWithTimeout (
141- group_id = "test" ,
142- timeout = TEST_TIMEOUT ,
143- ) as test :
144144 rollback_node_pool = node_pool .rollback .override (
145145 task_id = "rollback_node_pool"
146146 )(node_pool = cluster_info )
171171 )
172172
173173 chain (
174+ * startup .tasks ,
174175 rollback_node_pool ,
175176 wait_for_recovery ,
176177 verify_duration ,
Original file line number Diff line number Diff line change @@ -135,6 +135,10 @@ def get_current_time() -> TimeUtil:
135135 node_pool_selector = selector ,
136136 )
137137
138+ with TaskGroupWithTimeout (
139+ group_id = "test" ,
140+ timeout = TEST_TIMEOUT ,
141+ ) as test :
138142 startup = jobset .create_jobset_startup_tasks (
139143 node_pool = cluster_info ,
140144 jobset_config = jobset_config ,
@@ -143,10 +147,6 @@ def get_current_time() -> TimeUtil:
143147 workload_type = Workload .JAX_TPU_BENCHMARK ,
144148 )
145149
146- with TaskGroupWithTimeout (
147- group_id = "test" ,
148- timeout = TEST_TIMEOUT ,
149- ) as test :
150150 wait_for_jobset_uptime_data = (
151151 jobset .wait_for_jobset_uptime_data .override (
152152 task_id = "wait_for_jobset_uptime_data"
@@ -184,6 +184,7 @@ def get_current_time() -> TimeUtil:
184184 )
185185
186186 chain (
187+ * startup .tasks ,
187188 wait_for_jobset_uptime_data ,
188189 clean_up_workload ,
189190 jobset_clear_time ,
You can’t perform that action at this time.
0 commit comments