Skip to content

Commit 02f3d48

Browse files
authored
fix: Serialize v6e test execution to prevent timeouts (GoogleCloudPlatform#1080)
Run v6e tests one by one because running all v6e tests in parallel often leads to timeouts while waiting for resources.
1 parent 5676a58 commit 02f3d48

1 file changed

Lines changed: 31 additions & 8 deletions

File tree

dags/sparsity_diffusion_devx/project_bite_tpu_e2e.py

Lines changed: 31 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,12 @@
1616

1717
import datetime
1818
from airflow import models
19+
from airflow.utils.task_group import TaskGroup
20+
1921
from dags import composer_env
2022
from dags.common import test_owner
2123
from dags.common.vm_resource import TpuVersion, Zone, RuntimeVersion, Project
2224
from dags.sparsity_diffusion_devx.configs import project_bite_config as config
23-
from airflow.utils.task_group import TaskGroup
2425

2526

2627
# Run once a day at 6 pm UTC (11 am PST)
@@ -157,41 +158,63 @@
157158
group_id='bite_tpu_unittests', prefix_group_id=False
158159
) as bite_unittests:
159160
# Trillium (v6e) with JAX 0.5.3
160-
config.get_bite_tpu_unittests_config(
161+
jax_053_unittests_v6e_4 = config.get_bite_tpu_unittests_config(
161162
**trillium_conf,
162163
jax_version='0.5.3',
163164
**common,
164165
)
165166
# Trillium (v6e) with JAX 0.4.38
166-
config.get_bite_tpu_unittests_config(
167+
jax_0438_unittests_v6e_4 = config.get_bite_tpu_unittests_config(
167168
**trillium_conf,
168169
jax_version='0.4.38',
169170
**common,
170171
)
171172
# Trillium (v6e) with JAX nightly
172-
config.get_bite_tpu_unittests_config(
173+
jax_nightly_unittests_v6e_4 = config.get_bite_tpu_unittests_config(
173174
**trillium_conf,
174175
**common,
175176
)
176177
# V5P with JAX 0.5.3
177-
config.get_bite_tpu_unittests_config(
178+
jax_053_unittests_v5p_8 = config.get_bite_tpu_unittests_config(
178179
**v5p_conf,
179180
jax_version='0.5.3',
180181
**common,
181182
)
182183
# V5P with JAX 0.4.38
183-
config.get_bite_tpu_unittests_config(
184+
jax_0438_unittests_v5p_8 = config.get_bite_tpu_unittests_config(
184185
**v5p_conf,
185186
jax_version='0.4.38',
186187
**common,
187188
)
188189
# V5P with JAX nightly
189-
config.get_bite_tpu_unittests_config(
190+
jax_nightly_unittests_v5p_8 = config.get_bite_tpu_unittests_config(
190191
**v5p_conf,
191192
**common,
192193
)
193194
# V5E with JAX nightly
194-
config.get_bite_tpu_unittests_config(
195+
jax_nightly_unittests_v5e_8 = config.get_bite_tpu_unittests_config(
195196
**v5e_conf,
196197
**common,
197198
)
199+
200+
# Airflow uses >> for task chaining, which is pointless for pylint.
201+
# pylint: disable=pointless-statement
202+
[
203+
jax_053_unittests_v5p_8,
204+
jax_0438_unittests_v5p_8,
205+
jax_nightly_unittests_v5p_8,
206+
jax_nightly_unittests_v5e_8,
207+
jax_053_fuji_v5p_8,
208+
jax_main_fuji_v5p_8,
209+
]
210+
211+
# Running all v6e tests in parallel often leads to timeouts while waiting for resources.
212+
(
213+
jax_053_unittests_v6e_4
214+
>> jax_0438_unittests_v6e_4
215+
>> jax_nightly_unittests_v6e_4
216+
>> jax_053_fuji_v6e_8
217+
>> jax_main_fuji_v6e_8
218+
>> jax_pinned_053_fuji_v6e_8
219+
)
220+
# pylint: enable=pointless-statement

0 commit comments

Comments
 (0)