[rllib] - TaskPool.completed_prefetch() no longer returns stale object ids after an error (#7139)

This commit is contained in:
Adrian O'Grady
2020-02-13 22:30:44 -08:00
committed by GitHub
parent f3703bafa3
commit fe6ce714a0
3 changed files with 163 additions and 14 deletions
+11
View File
@@ -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"]
)
+14 -14
View File
@@ -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):
+138
View File
@@ -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)