Skip to content

Commit 3b6e4f0

Browse files
fix: Replace task chaining with chain for improved readability (GoogleCloudPlatform#1138)
This PR update the way to chain tasks in tpu_observability, from pylint “pointless-statement” into using from airflow.models.baseoperator import chain. Change 7 files for tpu_observability, including `interruption_validation_dag.py`, `jobset_ttr_rollback.py`, `multi_host_nodepool_rollback_dag.py`, `node_pool_status.py`, `node_pool_ttr_update_label.py`, `tpu_info_format_validation_dags.py` and `update_node_pool_label.py`.
1 parent 9e16b8a commit 3b6e4f0

7 files changed

Lines changed: 83 additions & 87 deletions

dags/tpu_observability/interruption_validation_dag.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from airflow import models
99
from airflow.decorators import task
1010
from airflow.exceptions import AirflowSkipException
11+
from airflow.models.baseoperator import chain
1112
from airflow.utils.task_group import TaskGroup
1213

1314
from dags.common import test_owner
@@ -567,10 +568,10 @@ def fetch_interruption_metric_records_task(
567568
log_records,
568569
)
569570

570-
(
571-
proper_time_range
572-
>> [metric_records, log_records]
573-
>> check_event_count
571+
chain(
572+
proper_time_range,
573+
[metric_records, log_records],
574+
check_event_count,
574575
)
575576

576577
return dag

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 10 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
import datetime
1818

1919
from airflow import models
20+
from airflow.models.baseoperator import chain
2021
from airflow.utils.trigger_rule import TriggerRule
2122
from airflow.utils.task_group import TaskGroup
2223

@@ -141,16 +142,13 @@
141142
setups=create_node_pool,
142143
)
143144

144-
# Airflow uses >> for task chaining, which is pointless for pylint.
145-
# pylint: disable=pointless-statement
146-
(
147-
cluster_info
148-
>> create_node_pool
149-
>> start_workload
150-
>> ensure_all_pods_running
151-
>> rollback_node_pool
152-
>> wait_for_metric_upload
153-
>> cleanup_workload
154-
>> cleanup_node_pool
145+
chain(
146+
cluster_info,
147+
create_node_pool,
148+
start_workload,
149+
ensure_all_pods_running,
150+
rollback_node_pool,
151+
wait_for_metric_upload,
152+
cleanup_workload,
153+
cleanup_node_pool,
155154
)
156-
# pylint: enable=pointless-statement

dags/tpu_observability/multi_host_nodepool_rollback_dag.py

Lines changed: 9 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
import datetime
2121

2222
from airflow import models
23+
from airflow.models.baseoperator import chain
2324
from airflow.utils.task_group import TaskGroup
2425
from airflow.utils.trigger_rule import TriggerRule
2526

@@ -117,15 +118,12 @@
117118
setups=create_node_pool,
118119
)
119120

120-
# Airflow uses >> for task chaining, which is pointless for pylint.
121-
# pylint: disable=pointless-statement
122-
(
123-
node_pool_info
124-
>> create_node_pool
125-
>> wait_node_pool_available
126-
>> rollback_node_pool
127-
>> wait_node_pool_unavailable
128-
>> wait_node_pool_recovered
129-
>> cleanup_node_pool
121+
chain(
122+
node_pool_info,
123+
create_node_pool,
124+
wait_node_pool_available,
125+
rollback_node_pool,
126+
wait_node_pool_unavailable,
127+
wait_node_pool_recovered,
128+
cleanup_node_pool,
130129
)
131-
# pylint: enable=pointless-statement

dags/tpu_observability/node_pool_status.py

Lines changed: 17 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818

1919
from airflow import models
2020
from airflow.decorators import task
21+
from airflow.models.baseoperator import chain
2122
from airflow.utils.task_group import TaskGroup
2223
from airflow.utils.trigger_rule import TriggerRule
2324

@@ -172,25 +173,22 @@ def generate_problematic_node_location(
172173
setups=create_problematic_node_pool_info,
173174
)
174175

175-
# Airflow uses >> for task chaining, which is pointless for pylint.
176-
# pylint: disable=pointless-statement
177-
normal_flow = (
178-
node_pool_info
179-
>> problematic_node_pool_info
180-
>> create_node_pool
181-
>> wait_for_provisioning
182-
>> wait_for_running
183-
>> delete_node
184-
>> wait_for_repair
185-
>> wait_for_recovered
186-
>> delete_node_pool
187-
>> wait_for_stopping
188-
>> cleanup_node_pool
176+
chain(
177+
node_pool_info,
178+
problematic_node_pool_info,
179+
create_node_pool,
180+
wait_for_provisioning,
181+
wait_for_running,
182+
delete_node,
183+
wait_for_repair,
184+
wait_for_recovered,
185+
delete_node_pool,
186+
wait_for_stopping,
187+
cleanup_node_pool,
189188
)
190189

191-
flow_for_error_state = (
192-
create_problematic_node_pool_info
193-
>> wait_for_error
194-
>> cleanup_wrong_node_pool
190+
chain(
191+
create_problematic_node_pool_info,
192+
wait_for_error,
193+
cleanup_wrong_node_pool,
195194
)
196-
# pylint: enable=pointless-statement

dags/tpu_observability/node_pool_ttr_update_label.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
import datetime
1818

1919
from airflow import models
20+
from airflow.models.baseoperator import chain
2021
from airflow.utils.task_group import TaskGroup
2122
from airflow.utils.trigger_rule import TriggerRule
2223

@@ -114,13 +115,13 @@
114115
setups=create_node_pool,
115116
)
116117

117-
_ = (
118-
node_pool_info
119-
>> create_node_pool
120-
>> wait_for_provisioning
121-
>> wait_for_running
122-
>> update_node_pool_label
123-
>> wait_for_recovered
124-
>> wait_for_ttr
125-
>> cleanup_node_pool
118+
chain(
119+
node_pool_info,
120+
create_node_pool,
121+
wait_for_provisioning,
122+
wait_for_running,
123+
update_node_pool_label,
124+
wait_for_recovered,
125+
wait_for_ttr,
126+
cleanup_node_pool,
126127
)

dags/tpu_observability/tpu_info_format_validation_dags.py

Lines changed: 23 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
from airflow import models
2929
from airflow.decorators import task
3030
from airflow.exceptions import AirflowFailException
31+
from airflow.models.baseoperator import chain
3132
from airflow.utils.task_group import TaskGroup
3233
from airflow.utils.trigger_rule import TriggerRule
3334

@@ -350,7 +351,10 @@ def generate_second_node_pool_name(
350351
tpu_accelerator_type="tpu-v6e-slice",
351352
tpu_topology="4x4",
352353
container_name="jax-tpu-worker",
353-
image="asia-northeast1-docker.pkg.dev/cienet-cmcs/yuna-docker/tpu-info:v0.5.1",
354+
image=(
355+
"asia-northeast1-docker.pkg.dev/cienet-cmcs/yuna-docker/"
356+
"tpu-info:v0.5.1"
357+
),
354358
tpu_cores_per_pod=4,
355359
)
356360

@@ -494,32 +498,30 @@ def generate_second_node_pool_name(
494498
setups=create_node_pool,
495499
)
496500

497-
# Airflow uses >> for task chaining, which is pointless for pylint.
498-
# pylint: disable=pointless-statement
499-
(
500-
verify_table_amount_task
501-
>> [
501+
chain(
502+
verify_table_amount_task,
503+
[
502504
validate_tpu_chips_metric,
503505
validate_runtime_metric,
504506
validate_tensorcore_metric,
505507
validate_latency_metric,
506-
]
508+
],
507509
)
508510

509511
[create_first_node_pool, create_second_node_pool]
510-
(cleanup_first_node_pool >> cleanup_second_node_pool)
511512

512-
(
513-
cluster_info
514-
>> cluster_info_2
515-
>> create_node_pool
516-
>> apply_time
517-
>> pod_names
518-
>> wait_for_job_start
519-
>> outputs_of_tpu_info
520-
>> output_of_tpu_info
521-
>> verification_group
522-
>> clean_up_workload
523-
>> cleanup_node_pool
513+
chain(cleanup_first_node_pool, cleanup_second_node_pool)
514+
515+
chain(
516+
cluster_info,
517+
cluster_info_2,
518+
create_node_pool,
519+
apply_time,
520+
pod_names,
521+
wait_for_job_start,
522+
outputs_of_tpu_info,
523+
output_of_tpu_info,
524+
verification_group,
525+
clean_up_workload,
526+
cleanup_node_pool,
524527
)
525-
# pylint: enable=pointless-statement

dags/tpu_observability/update_node_pool_label.py

Lines changed: 9 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
import datetime
2020

2121
from airflow import models
22+
from airflow.models.baseoperator import chain
2223
from airflow.utils.task_group import TaskGroup
2324
from airflow.utils.trigger_rule import TriggerRule
2425

@@ -109,15 +110,12 @@
109110
setups=[create_node_pool],
110111
)
111112

112-
# Airflow uses >> for task chaining, which is pointless for pylint.
113-
# pylint: disable=pointless-statement
114-
(
115-
node_pool_info
116-
>> create_node_pool
117-
>> wait_for_availability
118-
>> update_node_pool_label
119-
>> wait_for_unavailable
120-
>> wait_node_pool_recovered
121-
>> cleanup_node_pool
113+
chain(
114+
node_pool_info,
115+
create_node_pool,
116+
wait_for_availability,
117+
update_node_pool_label,
118+
wait_for_unavailable,
119+
wait_node_pool_recovered,
120+
cleanup_node_pool,
122121
)
123-
# pylint: enable=pointless-statement

0 commit comments

Comments
 (0)