[Serve] Batching in Worker Replica (#8709)

This commit is contained in:
Simon Mo
2020-06-09 11:29:16 -07:00
committed by GitHub
parent f007bfb4cf
commit 6c3062906f
20 changed files with 671 additions and 94 deletions
+111 -31
View File
@@ -1,20 +1,61 @@
import asyncio
import traceback
import inspect
from collections.abc import Iterable
from collections import defaultdict
from itertools import groupby
from operator import attrgetter
import time
import ray
from ray.async_compat import sync_to_async
from ray import serve
from ray.serve import context as serve_context
from ray.serve.context import FakeFlaskRequest
from collections import defaultdict
from ray.serve.utils import (parse_request_item, _get_logger)
from ray.serve.utils import (parse_request_item, _get_logger, chain_future,
unpack_future)
from ray.serve.exceptions import RayServeException
from ray.serve.metric import MetricClient
from ray.async_compat import sync_to_async
from ray.serve.config import BackendConfig
from ray.serve.router import Query
logger = _get_logger()
class WaitableQueue(asyncio.Queue):
async def wait_for_batch(self, num_items: int, timeout_s: float):
"""Wait up to num_items in the queue given timeout_s.
This method will block indefinitely for the first item. Therefore, it
guarantees to return at least one item.
"""
assert num_items >= 1
# Wait for the first value without timeout. We will return at least
# one item. Additionally this help the caller context switch on empty
# queue.
start_waiting = time.time()
batch = [
await self.get(),
]
# Adjust the timeout to account for the time waiting for first item.
time_remaining = timeout_s - (time.time() - start_waiting)
time_remaining = max(0, time_remaining)
# Wait for the remaining batch with the timeout
if num_items > 1:
done_set, not_done_set = await asyncio.wait(
[self.get() for _ in range(num_items - 1)],
timeout=time_remaining)
for task in done_set:
batch.append(task.result())
for task in not_done_set:
task.cancel()
return batch
def create_backend_worker(func_or_class):
"""Creates a worker class wrapping the provided function or class."""
@@ -30,8 +71,10 @@ def create_backend_worker(func_or_class):
backend_tag,
replica_tag,
init_args,
backend_config: BackendConfig,
instance_name=None):
serve.init(name=instance_name)
if is_function:
_callable = func_or_class
else:
@@ -42,11 +85,15 @@ def create_backend_worker(func_or_class):
metric_client = MetricClient(
metric_exporter, default_labels={"backend": backend_tag})
self.backend = RayServeWorker(backend_tag, replica_tag, _callable,
is_function, metric_client)
backend_config, is_function,
metric_client)
async def handle_request(self, request):
return await self.backend.handle_request(request)
def update_config(self, new_config: BackendConfig):
return self.backend.update_config(new_config)
def ready(self):
pass
@@ -75,13 +122,16 @@ def ensure_async(func):
class RayServeWorker:
"""Handles requests with the provided callable."""
def __init__(self, name, replica_tag, _callable, is_function,
metric_client):
def __init__(self, name, replica_tag, _callable,
backend_config: BackendConfig, is_function, metric_client):
self.name = name
self.replica_tag = replica_tag
self.callable = _callable
self.is_function = is_function
self.config = backend_config
self.query_queue = WaitableQueue()
self.metric_client = metric_client
self.request_counter = self.metric_client.new_counter(
"backend_request_counter",
@@ -101,6 +151,9 @@ class RayServeWorker:
self.restart_counter.labels(replica_tag=self.replica_tag).add()
self.loop_task = asyncio.get_event_loop().create_task(self.main_loop())
self.config_updated = asyncio.Event()
def get_runner_method(self, request_item):
method_name = request_item.call_method
if not hasattr(self.callable, method_name):
@@ -108,6 +161,8 @@ class RayServeWorker:
"which is specified in the request. "
"The available methods are {}".format(
method_name, dir(self.callable)))
if self.is_function:
return self.callable
return getattr(self.callable, method_name)
def has_positional_args(self, f):
@@ -124,6 +179,12 @@ class RayServeWorker:
return True
return False
def _reset_context(self):
# NOTE(simon): context management won't work in async mode because
# many concurrent queries might be running at the same time.
serve_context.web = None
serve_context.batch_size = None
async def invoke_single(self, request_item):
args, kwargs, is_web_context = parse_request_item(request_item)
serve_context.web = is_web_context
@@ -137,24 +198,12 @@ class RayServeWorker:
except Exception as e:
result = wrap_to_ray_error(e)
self.error_counter.add()
finally:
self._reset_context()
return result
async def invoke_batch(self, request_item_list):
# TODO(alind) : create no-http services. The enqueues
# from such services will always be TaskContext.Python.
# Assumption : all the requests in a bacth
# have same serve context.
# For batching kwargs are modified as follows -
# kwargs [Python Context] : key,val
# kwargs_list : key, [val1,val2, ... , valn]
# or
# args[Web Context] : val
# args_list : [val1,val2, ...... , valn]
# where n (current batch size) <= max_batch_size of a backend
arg_list = []
kwargs_list = defaultdict(list)
context_flags = set()
@@ -222,22 +271,53 @@ class RayServeWorker:
"results with length equal to the batch size"
".".format(batch_size, len(result_list)))
raise RayServeException(error_message)
self._reset_context()
return result_list
except Exception as e:
wrapped_exception = wrap_to_ray_error(e)
self.error_counter.add()
self._reset_context()
return [wrapped_exception for _ in range(batch_size)]
async def handle_request(self, request):
# check if work_item is a list or not
# if it is list: then batching supported
if not isinstance(request, list):
result = await self.invoke_single(request)
else:
result = await self.invoke_batch(request)
async def main_loop(self):
while True:
# NOTE(simon): There's an issue when user updated batch size and
# batch wait timeout during the execution, these values will not be
# updated until after the current iteration.
batch = await self.query_queue.wait_for_batch(
num_items=self.config.max_batch_size or 1,
timeout_s=self.config.batch_wait_timeout)
# re-assign to default values
serve_context.web = False
serve_context.batch_size = None
all_evaluated_futures = []
return result
if not self.config.accepts_batches:
query = batch[0]
evaluated = asyncio.ensure_future(self.invoke_single(query))
all_evaluated_futures = [evaluated]
chain_future(evaluated, query.async_future)
else:
get_call_method = attrgetter("call_method")
sorted_batch = sorted(batch, key=get_call_method)
for _, group in groupby(sorted_batch, key=get_call_method):
group = sorted(group)
evaluated = asyncio.ensure_future(self.invoke_batch(group))
all_evaluated_futures.append(evaluated)
result_futures = [q.async_future for q in group]
chain_future(
unpack_future(evaluated, len(group)), result_futures)
if self.config.is_blocking:
# We use asyncio.wait here so if the result is exception,
# it will not be raised.
await asyncio.wait(all_evaluated_futures)
def update_config(self, new_config: BackendConfig):
self.config = new_config
self.config_updated.set()
async def handle_request(self, request: Query):
assert not isinstance(request, list)
logger.debug("Worker {} got request {}".format(self.name, request))
request.async_future = asyncio.get_event_loop().create_future()
self.query_queue.put_nowait(request)
return await request.async_future