Skip to content

Commit e683b6e

Browse files
committed
move startup to test group
1 parent 1ee34c5 commit e683b6e

5 files changed

Lines changed: 25 additions & 20 deletions

File tree

dags/tpu_observability/jobset_ttr_drain_restart.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff 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,

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,10 @@
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,
@@ -133,10 +137,6 @@
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
)(
@@ -172,6 +172,7 @@
172172
)
173173

174174
chain(
175+
*startup.tasks,
175176
node_pool_resize_start_time,
176177
wait_for_recovery,
177178
verify_duration,

dags/tpu_observability/jobset_ttr_pod_delete.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -128,6 +128,10 @@
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,
@@ -136,10 +140,6 @@
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
)(
@@ -174,6 +174,7 @@
174174
)
175175

176176
chain(
177+
*startup.tasks,
177178
deletion_start_time,
178179
wait_for_recovery,
179180
verify_duration,

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,10 @@
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,
@@ -137,10 +141,6 @@
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)
@@ -171,6 +171,7 @@
171171
)
172172

173173
chain(
174+
*startup.tasks,
174175
rollback_node_pool,
175176
wait_for_recovery,
176177
verify_duration,

dags/tpu_observability/jobset_uptime_validation.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff 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,

0 commit comments

Comments
 (0)