Skip to content

Commit b6ed33c

Browse files
[llm][kv][11/N] Enable atomic selection & reservation broadcast to avoid herding (ray-project#65010)
Signed-off-by: Jeffrey Wang <jeffreywang@anyscale.com>
1 parent bc4dcc4 commit b6ed33c

18 files changed

Lines changed: 929 additions & 323 deletions

File tree

python/ray/llm/_internal/serve/constants.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,9 @@
6565
RAY_SERVE_LLM_ENABLE_DIRECT_STREAMING = (
6666
os.environ.get("RAY_SERVE_LLM_ENABLE_DIRECT_STREAMING", "0") == "1"
6767
)
68+
RAY_SERVE_LLM_ENABLE_DECODE_BLOCK_PROGRESS = (
69+
os.environ.get("RAY_SERVE_LLM_ENABLE_DECODE_BLOCK_PROGRESS", "0") == "1"
70+
)
6871

6972
MAX_NUM_STOPPING_SEQUENCES = int(os.getenv("RAYLLM_MAX_NUM_STOPPING_SEQUENCES", "8"))
7073
ENV_VARS_TO_PROPAGATE = {

python/ray/llm/_internal/serve/core/ingress/router.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -126,11 +126,13 @@ async def __init__(
126126
# tracked LLMServer deployment.
127127
from ray.llm._internal.serve.routing_policies.kv_aware.kv_token_tracker import ( # noqa: E501
128128
build_kv_token_tracker,
129+
get_llm_router_handle,
129130
)
130131

131132
self._kv_token_tracker = build_kv_token_tracker(
132133
llm_config, server.deployment_id
133134
)
135+
self._kv_token_tracker.start_reservation_broadcast(get_llm_router_handle())
134136
# Lazy import: this module pulls in vLLM's renderer;
135137
# keep it off the non-KV ingress import path.
136138
from ray.llm._internal.serve.routing_policies.kv_aware.tokenizer import (
@@ -206,6 +208,10 @@ async def on_lifecycle_events(self, batch):
206208
"""
207209
return await self._kv_token_tracker.on_lifecycle_events(batch)
208210

211+
async def on_reservations_created(self, batch):
212+
"""Ingress-facing intake for already-selected reservation bookings."""
213+
return await self._kv_token_tracker.on_reservations_created(batch)
214+
209215
async def _pick_replica(
210216
self,
211217
handle: DeploymentHandle,

python/ray/llm/_internal/serve/engines/vllm/vllm_engine.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,9 @@
2626
import ray
2727
from ray.llm._internal.common.callbacks.base import CallbackCtx
2828
from ray.llm._internal.common.utils.import_utils import try_import
29+
from ray.llm._internal.serve.constants import (
30+
RAY_SERVE_LLM_ENABLE_DECODE_BLOCK_PROGRESS,
31+
)
2932
from ray.llm._internal.serve.core.configs.llm_config import (
3033
DiskMultiplexConfig,
3134
LLMConfig,
@@ -571,7 +574,10 @@ def _start_async_llm_engine(
571574
# resolving it per request would block the engine's event loop.
572575
engine_cls = AsyncLLM
573576
if is_kv_aware(self.llm_config):
574-
engine_cls = enable_token_tracking(AsyncLLM)
577+
engine_cls = enable_token_tracking(
578+
AsyncLLM,
579+
report_decode_progress=RAY_SERVE_LLM_ENABLE_DECODE_BLOCK_PROGRESS,
580+
)
575581
engine_client = engine_cls(
576582
vllm_config=vllm_engine_config,
577583
executor_class=executor_class,

python/ray/llm/_internal/serve/routing_policies/kv_aware/kv_aware_router.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,18 @@
1919
logger = logging.getLogger(SERVE_LOGGER_NAME)
2020

2121

22+
def _get_expected_output_tokens(pending_request: PendingRequest) -> Optional[int]:
23+
"""The request's output cap from the routing payload, if present."""
24+
if not pending_request.args:
25+
return None
26+
payload = pending_request.args[0]
27+
for field in ("max_completion_tokens", "max_tokens"):
28+
value = getattr(payload, field, None)
29+
if isinstance(value, int) and value > 0:
30+
return value
31+
return None
32+
33+
2234
class KVAwareRouter(RequestRouter):
2335
"""Routes each request to the candidate that best balances expected KV-cache
2436
overlap against the worker's current prefill/decode load.
@@ -83,6 +95,7 @@ async def choose_replicas(
8395
pending_request.metadata.request_id,
8496
token_ids,
8597
list(worker_id_to_replica),
98+
_get_expected_output_tokens(pending_request),
8699
)
87100
return [[worker_id_to_replica[selection["worker_id"]]]]
88101

0 commit comments

Comments
 (0)