mirror of
https://github.com/wassname/ray.git
synced 2026-08-16 11:27:09 +08:00
[Serve] Batching in Worker Replica (#8709)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user