From fe6ce714a046e1b6a7c1e525208f1a8b27a51d67 Mon Sep 17 00:00:00 2001 From: Adrian O'Grady Date: Fri, 14 Feb 2020 06:30:44 +0000 Subject: [PATCH] [rllib] - TaskPool.completed_prefetch() no longer returns stale object ids after an error (#7139) --- rllib/BUILD | 11 +++ rllib/utils/actors.py | 28 +++--- rllib/utils/tests/test_taskpool.py | 138 +++++++++++++++++++++++++++++ 3 files changed, 163 insertions(+), 14 deletions(-) create mode 100644 rllib/utils/tests/test_taskpool.py diff --git a/rllib/BUILD b/rllib/BUILD index 073ea8cf6..11d09a2a2 100644 --- a/rllib/BUILD +++ b/rllib/BUILD @@ -48,3 +48,14 @@ py_test( size = "small", srcs = ["utils/schedules/tests/test_schedules.py"] ) + +# --------------------------------------- +# Utilities and Internals +# --------------------------------------- + +# TaskPool +py_test( + name = "test_taskpool", + size = "small", + srcs = ["utils/tests/test_taskpool.py"] +) diff --git a/rllib/utils/actors.py b/rllib/utils/actors.py index 676019110..073267be3 100644 --- a/rllib/utils/actors.py +++ b/rllib/utils/actors.py @@ -1,6 +1,7 @@ import logging import os import ray +from collections import deque logger = logging.getLogger(__name__) @@ -11,7 +12,7 @@ class TaskPool: def __init__(self): self._tasks = {} self._objects = {} - self._fetching = [] + self._fetching = deque() def add(self, worker, all_obj_ids): if isinstance(all_obj_ids, list): @@ -38,15 +39,11 @@ class TaskPool: for worker, obj_id in self.completed(blocking_wait=blocking_wait): self._fetching.append((worker, obj_id)) - remaining = [] - num_yielded = 0 - for worker, obj_id in self._fetching: - if num_yielded < max_yield: - yield (worker, obj_id) - num_yielded += 1 - else: - remaining.append((worker, obj_id)) - self._fetching = remaining + for _ in range(max_yield): + if not self._fetching: + break + + yield self._fetching.popleft() def reset_workers(self, workers): """Notify that some workers may be removed.""" @@ -54,11 +51,14 @@ class TaskPool: if ev not in workers: del self._tasks[obj_id] del self._objects[obj_id] - ok = [] - for ev, obj_id in self._fetching: + + # We want to keep the same deque reference so that we don't suffer from + # stale references in generators that are still in flight + for _ in range(len(self._fetching)): + ev, obj_id = self._fetching.popleft() if ev in workers: - ok.append((ev, obj_id)) - self._fetching = ok + # Re-queue items that are still valid + self._fetching.append((ev, obj_id)) @property def count(self): diff --git a/rllib/utils/tests/test_taskpool.py b/rllib/utils/tests/test_taskpool.py new file mode 100644 index 000000000..c6382206f --- /dev/null +++ b/rllib/utils/tests/test_taskpool.py @@ -0,0 +1,138 @@ +import unittest +from unittest.mock import patch + +import ray +from ray.rllib.utils.actors import TaskPool + + +def createMockWorkerAndObjectId(obj_id): + return ({obj_id: 1}, obj_id) + + +class TaskPoolTest(unittest.TestCase): + @patch("ray.wait") + def test_completed_prefetch_yieldsAllComplete(self, rayWaitMock): + task1 = createMockWorkerAndObjectId(1) + task2 = createMockWorkerAndObjectId(2) + # Return the second task as complete and the first as pending + rayWaitMock.return_value = ([2], [1]) + + pool = TaskPool() + pool.add(*task1) + pool.add(*task2) + + fetched = list(pool.completed_prefetch()) + self.assertListEqual(fetched, [task2]) + + @patch("ray.wait") + def test_completed_prefetch_yieldsAllCompleteUpToDefaultLimit( + self, rayWaitMock): + # Load the pool with 1000 tasks, mock them all as complete and then + # check that the first call to completed_prefetch only yields 999 + # items and the second call yields the final one + pool = TaskPool() + for i in range(1000): + task = createMockWorkerAndObjectId(i) + pool.add(*task) + + rayWaitMock.return_value = (list(range(1000)), []) + + # For this test, we're only checking the object ids + fetched = [pair[1] for pair in pool.completed_prefetch()] + self.assertListEqual(fetched, list(range(999))) + + # Finally, check the next iteration returns the final taks + fetched = [pair[1] for pair in pool.completed_prefetch()] + self.assertListEqual(fetched, [999]) + + @patch("ray.wait") + def test_completed_prefetch_yieldsAllCompleteUpToSpecifiedLimit( + self, rayWaitMock): + # Load the pool with 1000 tasks, mock them all as complete and then + # check that the first call to completed_prefetch only yield 999 items + # and the second call yields the final one + pool = TaskPool() + for i in range(1000): + task = createMockWorkerAndObjectId(i) + pool.add(*task) + + rayWaitMock.return_value = (list(range(1000)), []) + + # Verify that only the first 500 tasks are returned, this should leave + # some tasks in the _fetching deque for later + fetched = [pair[1] for pair in pool.completed_prefetch(max_yield=500)] + self.assertListEqual(fetched, list(range(500))) + + # Finally, check the next iteration returns the remaining tasks + fetched = [pair[1] for pair in pool.completed_prefetch()] + self.assertListEqual(fetched, list(range(500, 1000))) + + @patch("ray.wait") + def test_completed_prefetch_yieldsRemainingIfIterationStops( + self, rayWaitMock): + # Test for issue #7106 + # In versions of Ray up to 0.8.1, if the pre-fetch generator failed to + # run to completion, then the TaskPool would fail to clear up already + # fetched tasks resulting in stale object ids being returned + pool = TaskPool() + for i in range(10): + task = createMockWorkerAndObjectId(i) + pool.add(*task) + + rayWaitMock.return_value = (list(range(10)), []) + + # This should fetch just the first item in the list + try: + for _ in pool.completed_prefetch(): + # Simulate a worker failure returned by ray.get() + raise ray.exceptions.RayError + except ray.exceptions.RayError: + pass + + # This fetch should return the remaining pre-fetched tasks + fetched = [pair[1] for pair in pool.completed_prefetch()] + self.assertListEqual(fetched, list(range(1, 10))) + + @patch("ray.wait") + def test_reset_workers_pendingFetchesFromFailedWorkersRemoved( + self, rayWaitMock): + pool = TaskPool() + # We need to hold onto the tasks for this test so that we can fail a + # specific worker + tasks = [] + + for i in range(10): + task = createMockWorkerAndObjectId(i) + pool.add(*task) + tasks.append(task) + + # Simulate only some of the work being complete and fetch a couple of + # tasks in order to fill the fetching queue + rayWaitMock.return_value = ([0, 1, 2, 3, 4, 5], [6, 7, 8, 9]) + fetched = [pair[1] for pair in pool.completed_prefetch(max_yield=2)] + + # As we still have some pending tasks, we need to update the + # completion states to remove the completed tasks + rayWaitMock.return_value = ([], [6, 7, 8, 9]) + + pool.reset_workers([ + tasks[0][0], + tasks[1][0], + tasks[2][0], + tasks[3][0], + # OH NO! WORKER 4 HAS CRASHED! + tasks[5][0], + tasks[6][0], + tasks[7][0], + tasks[8][0], + tasks[9][0] + ]) + + # Fetch the remaining tasks which should already be in the _fetching + # queue + fetched = [pair[1] for pair in pool.completed_prefetch()] + self.assertListEqual(fetched, [2, 3, 5]) + + +if __name__ == "__main__": + unittest.main(verbosity=2)