Skip to content

Commit 2b254c9

Browse files
committed
Keep track of the futures
1 parent ea56fbf commit 2b254c9

3 files changed

Lines changed: 64 additions & 39 deletions

File tree

src/executorlib/standalone/batched.py

Lines changed: 12 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -3,35 +3,34 @@
33

44
def batched_futures(
55
lst: list[Future], nested_skip_lst: list[Future[list]], n: int
6-
) -> list[list] | BaseException:
6+
) -> tuple[bool, list[Future]]:
77
"""
88
Batch n completed future objects. If the number of completed futures is smaller than n and the end of the batch is
99
not reached yet, then an empty list is returned. If n future objects are done, which are not included in the skip_set
1010
then they are returned as batch.
1111
1212
Args:
1313
lst (list): list of all future objects
14-
nested_skip_lst (list): nest list of individual results already assigned to previous batches
14+
nested_skip_lst (list): list of future objects, which contain the list of future objects ids which should be skipped for the batch
1515
n (int): batch size
1616
1717
Returns:
1818
list: results of the batched futures
1919
"""
20-
skip_set = {id(item) for f in nested_skip_lst for item in f.result()}
20+
skip_set = {fid for f in nested_skip_lst for fid in f.result()}
2121

2222
done_lst = []
2323
failed_lst = []
2424
n_expected = min(n, len(lst) - len(skip_set))
2525
for v in lst:
26-
if v.done():
27-
excp = v.exception()
28-
if excp is not None:
29-
failed_lst.append(excp)
30-
elif id(v.result()) not in skip_set:
31-
done_lst.append(v.result())
26+
if id(v) not in skip_set and v.done():
27+
if v.exception() is not None:
28+
failed_lst.append(v)
29+
elif id(v) not in skip_set and v.done():
30+
done_lst.append(v)
3231
if len(done_lst) == n_expected:
33-
return done_lst
34-
if len(failed_lst) == len(lst) - len(skip_set) and len(failed_lst) > 0:
35-
return failed_lst[0] # raise the exception only after all futures have failed
32+
return True, done_lst
33+
if (len(lst) - len(skip_set)) == len(failed_lst):
34+
return False, failed_lst[:n_expected] # raise the exception only after all futures have failed
3635
else:
37-
return []
36+
return True, []

src/executorlib/task_scheduler/interactive/dependency.py

Lines changed: 11 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -177,17 +177,19 @@ def batched(
177177
future_lst: list[Future] = []
178178
for _ in range(len(iterable) // n + (1 if len(iterable) % n > 0 else 0)):
179179
f: Future = Future()
180+
f_skip = Future()
180181
if self._future_queue is not None:
181182
self._future_queue.put(
182183
{
183184
"fn": "batched",
184185
"args": (),
185186
"kwargs": {"lst": iterable, "n": n, "skip_lst": skip_lst},
186187
"future": f,
188+
"future_skip": f_skip,
187189
"resource_dict": {},
188190
}
189191
)
190-
skip_lst = skip_lst.copy() + [f] # be careful
192+
skip_lst = skip_lst.copy() + [f_skip] # be careful
191193
future_lst.append(f)
192194

193195
return future_lst
@@ -330,7 +332,7 @@ def _update_waiting_task(
330332
wait_tmp_lst = []
331333
for task_wait_dict in wait_lst:
332334
exception_lst = get_exception_lst(future_lst=task_wait_dict["future_lst"])
333-
if len(exception_lst) > 0:
335+
if len(exception_lst) > 0 and task_wait_dict["fn"] != "batched":
334336
task_wait_dict["future"].set_exception(exception_lst[0])
335337
elif task_wait_dict["fn"] != "batched" and all(
336338
future.done() for future in task_wait_dict["future_lst"]
@@ -343,17 +345,19 @@ def _update_waiting_task(
343345
elif task_wait_dict["fn"] == "batched" and all(
344346
future.done() for future in task_wait_dict["kwargs"]["skip_lst"]
345347
):
346-
done_lst = batched_futures(
348+
success, done_lst = batched_futures(
347349
lst=task_wait_dict["kwargs"]["lst"],
348350
n=task_wait_dict["kwargs"]["n"],
349351
nested_skip_lst=task_wait_dict["kwargs"]["skip_lst"],
350352
)
351-
if isinstance(done_lst, list) and len(done_lst) == 0:
353+
if success and len(done_lst) == 0:
352354
wait_tmp_lst.append(task_wait_dict)
353-
elif isinstance(done_lst, list) and len(done_lst) > 0:
354-
task_wait_dict["future"].set_result(done_lst)
355+
elif success and len(done_lst) > 0:
356+
task_wait_dict["future"].set_result([f.result() for f in done_lst])
357+
task_wait_dict["future_skip"].set_result([id(f) for f in done_lst])
355358
else:
356-
task_wait_dict["future"].set_exception(done_lst)
359+
task_wait_dict["future"].set_exception(done_lst[0].exception())
360+
task_wait_dict["future_skip"].set_result([id(f) for f in done_lst])
357361
else:
358362
wait_tmp_lst.append(task_wait_dict)
359363
if len(wait_lst) == len(wait_tmp_lst):

tests/unit/standalone/test_batched.py

Lines changed: 41 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -12,13 +12,21 @@ def test_batched_futures(self):
1212
f.set_result(i)
1313
lst.append(f)
1414
batched_lst = [Future(), Future(), Future()]
15-
batched_lst[0].set_result([0, 1, 2])
16-
batched_lst[1].set_result([3, 4, 5])
17-
batched_lst[2].set_result([6, 7, 8])
18-
self.assertEqual(batched_futures(lst=lst, n=3, nested_skip_lst=set()), [0, 1, 2])
19-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3), [3, 4, 5])
20-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst[:2], n=3), [6, 7, 8])
21-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst, n=3), [9])
15+
batched_lst[0].set_result([id(lst[0]), id(lst[1]), id(lst[2])])
16+
batched_lst[1].set_result([id(lst[3]), id(lst[4]), id(lst[5])])
17+
batched_lst[2].set_result([id(lst[6]), id(lst[7]), id(lst[8])])
18+
success, done_lst = batched_futures(lst=lst, n=3, nested_skip_lst=set())
19+
self.assertTrue(success)
20+
self.assertEqual([f.result() for f in done_lst], [0, 1, 2])
21+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3)
22+
self.assertTrue(success)
23+
self.assertEqual([f.result() for f in done_lst], [3, 4, 5])
24+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst[:2], n=3)
25+
self.assertTrue(success)
26+
self.assertEqual([f.result() for f in done_lst], [6, 7, 8])
27+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst, n=3)
28+
self.assertTrue(success)
29+
self.assertEqual([f.result() for f in done_lst], [9])
2230

2331
def test_batched_futures_duplicated(self):
2432
lst = []
@@ -28,12 +36,18 @@ def test_batched_futures_duplicated(self):
2836
f.set_result(i)
2937
lst.append(f)
3038
batched_lst = [Future(), Future(), Future()]
31-
batched_lst[0].set_result([1, 1, 1])
32-
batched_lst[1].set_result([2, 2, 2])
33-
batched_lst[2].set_result([3, 3, 3])
34-
self.assertEqual(batched_futures(lst=lst, n=3, nested_skip_lst=set()), [1, 1, 1])
35-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3), [2, 2, 2])
36-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst[:2], n=3), [3, 3, 3])
39+
batched_lst[0].set_result([id(lst[0]), id(lst[1]), id(lst[2])])
40+
batched_lst[1].set_result([id(lst[3]), id(lst[4]), id(lst[5])])
41+
batched_lst[2].set_result([id(lst[6]), id(lst[7]), id(lst[8])])
42+
success, done_lst = batched_futures(lst=lst, n=3, nested_skip_lst=set())
43+
self.assertTrue(success)
44+
self.assertEqual([f.result() for f in done_lst], [1, 1, 1])
45+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3)
46+
self.assertTrue(success)
47+
self.assertEqual([f.result() for f in done_lst], [2, 2, 2])
48+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst[:2], n=3)
49+
self.assertTrue(success)
50+
self.assertEqual([f.result() for f in done_lst], [3, 3, 3])
3751

3852
def test_batched_futures(self):
3953
lst = []
@@ -45,16 +59,24 @@ def test_batched_futures(self):
4559
f.set_result(i)
4660
lst.append(f)
4761
batched_lst = [Future(), Future()]
48-
batched_lst[0].set_result([1, 2, 4])
49-
batched_lst[1].set_result([5, 7, 8])
50-
self.assertEqual(batched_futures(lst=lst, n=3, nested_skip_lst=set()), [1, 2, 4])
51-
self.assertEqual(batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3), [5, 7, 8])
62+
batched_lst[0].set_result([id(lst[1]), id(lst[2]), id(lst[4])])
63+
batched_lst[1].set_result([id(lst[5]), id(lst[7]), id(lst[8])])
64+
success, done_lst = batched_futures(lst=lst, n=3, nested_skip_lst=set())
65+
self.assertTrue(success)
66+
self.assertEqual([f.result() for f in done_lst], [1, 2, 4])
67+
success, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst[:1], n=3)
68+
self.assertTrue(success)
69+
self.assertEqual([f.result() for f in done_lst], [5, 7, 8])
70+
succss, done_lst = batched_futures(lst=lst, nested_skip_lst=batched_lst, n=3)
71+
self.assertFalse(succss)
5272
with self.assertRaises(ValueError):
53-
raise batched_futures(lst=lst, nested_skip_lst=batched_lst, n=3)
73+
raise done_lst[0].exception()
5474

5575
def test_batched_futures_not_finished(self):
5676
lst = []
5777
for _ in list(range(10)):
5878
f = Future()
5979
lst.append(f)
60-
self.assertEqual(batched_futures(lst=lst, n=3, nested_skip_lst=set()), [])
80+
success, done_lst = batched_futures(lst=lst, n=3, nested_skip_lst=set())
81+
self.assertTrue(success)
82+
self.assertEqual(done_lst, [])

0 commit comments

Comments
 (0)