Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions components/src/dynamo/router/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,7 +137,9 @@ async def generate(self, request):
}
yield llm_engine_output

async def best_worker_id(self, token_ids, router_config_override=None):
async def best_worker_id(
self, token_ids, router_config_override=None, allowed_worker_ids=None
):
"""
Get the best worker ID for a given set of tokens without actually routing.

Expand All @@ -150,7 +152,9 @@ async def best_worker_id(self, token_ids, router_config_override=None):
raise RuntimeError("Router not initialized")

(worker_id, _dp_rank, _overlap_blocks) = await self.kv_router.best_worker(
token_ids, router_config_override
token_ids,
router_config_override,
allowed_worker_ids=allowed_worker_ids,
)

yield worker_id
Expand Down
131 changes: 126 additions & 5 deletions lib/bindings/python/rust/llm/kv.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
// SPDX-License-Identifier: Apache-2.0

use pythonize::{depythonize, pythonize};
use std::collections::HashMap;
use std::collections::{HashMap, HashSet};
use std::ffi::OsString;
use std::sync::Arc;
use std::sync::atomic::AtomicU32;
Expand Down Expand Up @@ -36,6 +36,10 @@ fn depythonize_block_mm_infos(obj: &Bound<'_, PyAny>) -> PyResult<Vec<Option<Blo
depythonize(obj).map_err(to_pyerr)
}

fn allowed_worker_set(allowed_worker_ids: Option<Vec<WorkerId>>) -> Option<HashSet<WorkerId>> {
allowed_worker_ids.map(|ids| ids.into_iter().collect())
}

#[cfg(feature = "kv-indexer")]
#[derive(Parser)]
#[command(
Expand Down Expand Up @@ -1094,7 +1098,7 @@ impl KvRouter {
}

#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (token_ids, router_config_override=None, request_id=None, update_indexer=false, block_mm_infos=None, lora_name=None, routing_constraints=None))]
#[pyo3(signature = (token_ids, router_config_override=None, request_id=None, update_indexer=false, block_mm_infos=None, lora_name=None, routing_constraints=None, allowed_worker_ids=None))]
fn best_worker<'p>(
&self,
py: Python<'p>,
Expand All @@ -1105,6 +1109,7 @@ impl KvRouter {
block_mm_infos: Option<PyObject>,
lora_name: Option<String>,
routing_constraints: Option<RoutingConstraints>,
allowed_worker_ids: Option<Vec<WorkerId>>,
) -> PyResult<Bound<'p, PyAny>> {
let router_config_override = if let Some(obj) = router_config_override {
let override_config: RouterConfigOverride =
Expand All @@ -1117,11 +1122,16 @@ impl KvRouter {
let block_mm_infos = block_mm_infos
.map(|obj| depythonize_block_mm_infos(obj.bind(py)))
.transpose()?;
let allowed_worker_ids = allowed_worker_set(allowed_worker_ids);

let chooser = self.inner.chooser.clone();
let update_states = request_id.is_some();

pyo3_async_runtimes::tokio::future_into_py(py, async move {
if let Some(worker_ids) = allowed_worker_ids.as_ref() {
chooser.register_workers(worker_ids);
}

let outcome = chooser
.find_best_match_details(
request_id.as_deref(),
Expand All @@ -1134,7 +1144,7 @@ impl KvRouter {
0.0,
None,
None,
None, // allowed_worker_ids: pass via RoutingHints in PreprocessedRequest path
allowed_worker_ids,
routing_constraints.map(Into::into).unwrap_or_default(),
)
.await
Expand Down Expand Up @@ -1195,6 +1205,108 @@ impl KvRouter {
})
}

#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (token_ids, router_config_override=None, allowed_worker_ids=None, block_mm_infos=None, lora_name=None))]
fn rank_workers<'p>(
&self,
py: Python<'p>,
token_ids: Vec<u32>,
router_config_override: Option<PyObject>,
allowed_worker_ids: Option<Vec<WorkerId>>,
block_mm_infos: Option<PyObject>,
lora_name: Option<String>,
) -> PyResult<Bound<'p, PyAny>> {
let router_config_override = if let Some(obj) = router_config_override {
let override_config: RouterConfigOverride =
depythonize(obj.bind(py)).map_err(to_pyerr)?;
Some(override_config)
} else {
None
};

let block_mm_infos = block_mm_infos
.map(|obj| depythonize_block_mm_infos(obj.bind(py)))
.transpose()?;
let allowed_worker_ids = allowed_worker_set(allowed_worker_ids);

let chooser = self.inner.chooser.clone();

pyo3_async_runtimes::tokio::future_into_py(py, async move {
let ranked = chooser
.rank_workers(
&token_ids,
block_mm_infos.as_deref(),
router_config_override.as_ref(),
lora_name,
allowed_worker_ids,
)
.await
.map_err(to_pyerr)?;

Python::with_gil(|py| {
pythonize(py, &ranked)
.map(|obj| obj.unbind())
.map_err(to_pyerr)
})
})
}

#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (request_id, token_ids, worker_id, dp_rank=0, overlap_blocks=0, expected_output_tokens=None, router_config_override=None, block_mm_infos=None, lora_name=None))]
fn add_request<'p>(
&self,
py: Python<'p>,
request_id: String,
token_ids: Vec<u32>,
worker_id: WorkerId,
dp_rank: DpRank,
overlap_blocks: u32,
expected_output_tokens: Option<u32>,
router_config_override: Option<PyObject>,
block_mm_infos: Option<PyObject>,
lora_name: Option<String>,
) -> PyResult<Bound<'p, PyAny>> {
let router_config_override = if let Some(obj) = router_config_override {
let override_config: RouterConfigOverride =
depythonize(obj.bind(py)).map_err(to_pyerr)?;
Some(override_config)
} else {
None
};

let block_mm_infos = block_mm_infos
.map(|obj| depythonize_block_mm_infos(obj.bind(py)))
.transpose()?;

let chooser = self.inner.chooser.clone();
let worker = WorkerWithDpRank::new(worker_id, dp_rank);

pyo3_async_runtimes::tokio::future_into_py(py, async move {
let cached_tokens = (overlap_blocks as usize)
.saturating_mul(chooser.block_size() as usize)
.min(token_ids.len());
chooser
.add_request(
request_id,
&token_ids,
block_mm_infos.as_deref(),
cached_tokens,
expected_output_tokens,
worker,
lora_name,
router_config_override.as_ref(),
)
.await;
Ok::<(), pyo3::PyErr>(())
})
}

#[pyo3(signature = (worker_ids))]
fn register_workers(&self, worker_ids: Vec<WorkerId>) {
let worker_ids: HashSet<WorkerId> = worker_ids.into_iter().collect();
self.inner.chooser.register_workers(&worker_ids);
}

/// Mark prefill as completed for a request
fn mark_prefill_complete<'p>(
&self,
Expand Down Expand Up @@ -1222,21 +1334,27 @@ impl KvRouter {
})
}

#[pyo3(signature = (token_ids, block_mm_infos=None, lora_name=None))]
#[pyo3(signature = (token_ids, block_mm_infos=None, lora_name=None, allowed_worker_ids=None))]
fn get_potential_loads<'p>(
&self,
py: Python<'p>,
token_ids: Vec<u32>,
block_mm_infos: Option<PyObject>,
lora_name: Option<String>,
allowed_worker_ids: Option<Vec<WorkerId>>,
) -> PyResult<Bound<'p, PyAny>> {
let block_mm_infos = block_mm_infos
.map(|obj| depythonize_block_mm_infos(obj.bind(py)))
.transpose()?;
let allowed_worker_ids = allowed_worker_set(allowed_worker_ids);
let chooser = self.inner.chooser.clone();

pyo3_async_runtimes::tokio::future_into_py(py, async move {
let loads = chooser
if let Some(worker_ids) = allowed_worker_ids.as_ref() {
chooser.register_workers(worker_ids);
}

let mut loads = chooser
.get_potential_loads(
&token_ids,
None,
Expand All @@ -1245,6 +1363,9 @@ impl KvRouter {
)
.await
.map_err(to_pyerr)?;
if let Some(allowed_worker_ids) = allowed_worker_ids {
loads.retain(|load| allowed_worker_ids.contains(&load.worker_id));
}

// Return loads without aggregation - each (worker_id, dp_rank) pair is a separate entry
// Use pythonize to convert Vec<PotentialLoad> to Python list of dicts
Expand Down
72 changes: 72 additions & 0 deletions lib/bindings/python/src/dynamo/_core.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -2362,6 +2362,7 @@ class KvRouter:
block_mm_infos: Optional[List[Optional[Dict[str, Any]]]] = None,
lora_name: Optional[str] = None,
routing_constraints: Optional[RoutingConstraints] = None,
allowed_worker_ids: Optional[List[int]] = None,
) -> Tuple[int, int, int]:
"""
Find the best matching worker for the given tokens.
Expand All @@ -2379,6 +2380,11 @@ class KvRouter:
block_mm_infos: Optional block-level multimodal metadata aligned to request
blocks. When provided, this is used in block hash computation
to enable MM-aware worker selection.
routing_constraints: Optional topology or taint constraints to apply while
selecting a worker.
allowed_worker_ids: Optional worker IDs to consider. If provided, routing is
restricted to this set and the IDs are lazily
registered for external/Ray-owned scoring.

Returns:
A tuple of (worker_id, dp_rank, overlap_blocks) where:
Expand All @@ -2388,11 +2394,76 @@ class KvRouter:
"""
...

async def rank_workers(
self,
token_ids: List[int],
router_config_override: Optional[JsonLike] = None,
allowed_worker_ids: Optional[List[int]] = None,
block_mm_infos: Optional[List[Optional[Dict[str, Any]]]] = None,
lora_name: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""
Rank candidate workers for the given tokens without booking request state.

Args:
token_ids: List of token IDs to rank workers for.
router_config_override: Optional router configuration override.
allowed_worker_ids: Optional worker IDs to consider. If provided,
ranking is restricted to this set and the IDs are
lazily registered for external/Ray-owned scoring.
block_mm_infos: Optional block-level multimodal metadata aligned to
request blocks.
lora_name: Optional LoRA adapter name used in block hash computation.

Returns:
A list of dictionaries sorted from best to worst. Each dictionary contains:
- worker_id: The worker ID
- dp_rank: The data parallel rank
- overlap_blocks: Number of matching cached blocks
- potential_prefill_tokens: Estimated prefill tokens if routed here
- potential_decode_blocks: Estimated active decode blocks
- score: Final router score, lower is better
"""
...

async def add_request(
self,
request_id: str,
token_ids: List[int],
worker_id: int,
dp_rank: int = 0,
overlap_blocks: int = 0,
expected_output_tokens: Optional[int] = None,
router_config_override: Optional[JsonLike] = None,
block_mm_infos: Optional[List[Optional[Dict[str, Any]]]] = None,
lora_name: Optional[str] = None,
) -> None:
"""
Book a request in router load state after an external dispatcher accepts it.

This is the explicit lifecycle counterpart to query-only `best_worker()` or
`rank_workers()` calls. It should be paired with `mark_prefill_complete()`
and `free()`.
"""
...

def register_workers(self, worker_ids: List[int]) -> None:
"""
Add externally-known workers to the router's active-load tracker.

This is additive and does not remove workers absent from `worker_ids`.
Ray or another external router remains the source of feasible replica
membership. If Dynamo already knows a worker's DP config it is reused;
otherwise the worker is treated as a single-rank cold worker.
"""
...

async def get_potential_loads(
self,
token_ids: List[int],
block_mm_infos: Optional[List[Optional[Dict[str, Any]]]] = None,
lora_name: Optional[str] = None,
allowed_worker_ids: Optional[List[int]] = None,
) -> List[Dict[str, int]]:
"""
Get potential prefill and decode loads for all workers.
Expand All @@ -2402,6 +2473,7 @@ class KvRouter:
block_mm_infos: Optional block-level multimodal metadata aligned to request
blocks. When provided, this is used in hash computation
for MM-aware potential-load estimation.
allowed_worker_ids: Optional worker IDs to include in the returned load list.

Returns:
A list of dictionaries, each containing:
Expand Down
Loading
Loading