Skip to content
Merged
50 changes: 47 additions & 3 deletions components/src/dynamo/vllm/handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from typing import Any, AsyncGenerator, Dict, Final

import torch
from vllm.config import VllmConfig
from vllm.inputs import EmbedsPrompt, TextPrompt, TokensPrompt
from vllm.lora.request import LoRARequest
from vllm.outputs import RequestOutput
Expand Down Expand Up @@ -260,6 +261,32 @@ def build_sampling_params_openai(
return sampling_params


def get_dp_range_for_worker(vllm_config: VllmConfig) -> range:
"""
Get the global DP rank range that this worker is responsible for based on vLLM config.
Note that the 'vllm_config' is normalized so the load balancing flags are set properly.
The return value is in the format of (start_dp_rank, managed_dp_size)."""
if vllm_config.parallel_config.data_parallel_external_lb:
# external load balancing, each worker is responsible for exactly 1 rank
return (vllm_config.parallel_config.data_parallel_rank, 1)
elif vllm_config.parallel_config.data_parallel_hybrid_lb:
# hybrid load balancing, each worker is responsible for a subset of local ranks
return (
vllm_config.parallel_config.data_parallel_rank,
vllm_config.parallel_config.data_parallel_size_local,
)
else:
# internal load balancing, the worker is responsible for all DP ranks
Comment thread
GuanLuo marked this conversation as resolved.
logger.warning(
"vLLM selects internal DP load balancing. If you are launching multiple workers for DP deployment,"
" hybrid or external load balancing is recommended."
)
return (
vllm_config.parallel_config.data_parallel_rank,
vllm_config.parallel_config.data_parallel_size,
)


class BaseWorkerHandler(ABC):
"""
Request handler for the generate and clear_kv_blocks endpoints.
Expand Down Expand Up @@ -302,6 +329,8 @@ def __init__(

self.use_vllm_tokenizer = use_vllm_tokenizer

self.dp_range = get_dp_range_for_worker(self.engine_client.vllm_config)

# Initialize InputParamManager for text-in-text-out mode
tokenizer = None
if use_vllm_tokenizer and hasattr(engine, "tokenizer"):
Expand Down Expand Up @@ -463,6 +492,21 @@ def add_temp_dir(self, temp_dir: tempfile.TemporaryDirectory) -> None:
if temp_dir is not None:
self.temp_dirs.append(temp_dir)

def _to_local_dp_rank(self, dp_rank: int | None) -> int | None:
"""Convert global DP rank to local DP rank based on engine config."""
if dp_rank is None:
return None
if dp_rank < self.dp_range[0] or dp_rank >= self.dp_range[0] + self.dp_range[1]:
logger.warning(
f"Received DP rank {dp_rank} is out of range [{self.dp_range[0]} - {self.dp_range[0] + self.dp_range[1]}), fallback to vLLM internal DP selection"
)
return None
local_dp_rank = (dp_rank - self.dp_range[0]) % self.dp_range[1]
Comment thread
GuanLuo marked this conversation as resolved.
logger.debug(
f"Converted global DP rank {dp_rank} to local DP rank {local_dp_rank}"
)
return local_dp_rank

def _resolve_lora_request(self, model_name: str | None) -> LoRARequest | None:
"""Return a LoRARequest if model_name is a loaded adapter, else None."""
if model_name and (lora := self.loaded_loras.get(model_name)):
Expand Down Expand Up @@ -1317,7 +1361,7 @@ async def _generate_token_mode(self, request, context, request_id):
f"Decode request {request_id} has no LoRA specified (model: {model_name})"
)
routing = request.get("routing") or {}
dp_rank = routing.get("dp_rank")
dp_rank = self._to_local_dp_rank(routing.get("dp_rank"))
priority = routing.get("priority", 0)

trace_headers = build_trace_headers(context)
Expand Down Expand Up @@ -1364,7 +1408,7 @@ async def _generate_text_mode(self, request, context, request_id):
)

routing = request.get("routing") or {}
dp_rank = routing.get("dp_rank")
dp_rank = self._to_local_dp_rank(routing.get("dp_rank"))
priority = routing.get("priority", 0)
openai_request_id = request.get("id") or request.get("request_id", request_id)
previous_text = ""
Expand Down Expand Up @@ -1525,7 +1569,7 @@ async def _generate_token_mode(self, request, context, request_id):
)

routing = request.get("routing") or {}
dp_rank = routing.get("dp_rank")
dp_rank = self._to_local_dp_rank(routing.get("dp_rank"))
priority = routing.get("priority", 0)

trace_headers = build_trace_headers(context)
Expand Down
43 changes: 8 additions & 35 deletions components/src/dynamo/vllm/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,7 @@
from .args import Config, _uses_dynamo_connector, parse_args
from .checkpoint_restore import get_checkpoint_config
from .constants import DisaggregationMode
from .handlers import DecodeWorkerHandler, PrefillWorkerHandler
from .handlers import DecodeWorkerHandler, PrefillWorkerHandler, get_dp_range_for_worker
from .health_check import (
VllmHealthCheckPayload,
VllmOmniHealthCheckPayload,
Expand All @@ -69,19 +69,6 @@
CHECKPOINT_SLEEP_MODE_LEVEL = 1


async def _handle_non_leader_node(dp_rank: int) -> None:
Comment thread
alec-flowers marked this conversation as resolved.
"""
Handle non-leader node (data_parallel_rank >= 1) in multi-node deployments.
Non-leader nodes run vLLM workers but don't serve Dynamo endpoints.
"""
logger.info(
f"Non-leader node detected (data_parallel_rank={dp_rank}). "
"Skipping endpoint serving."
)
# Wait indefinitely - process terminated via signal handlers
await asyncio.Event().wait()


def build_headless_namespace(config: Config) -> argparse.Namespace:
"""Build an argparse Namespace from engine_args for vLLM's run_headless().

Expand Down Expand Up @@ -339,11 +326,12 @@ def setup_kv_event_publisher(
)
return None

# Get data_parallel_size to create publishers for all dp_ranks
data_parallel_size = getattr(vllm_config.parallel_config, "data_parallel_size", 1)
# Get DP rank range managed by this worker to create publishers for corresponding dp_ranks,
# all served workers should cover all ranks.
dp_start, dp_size = get_dp_range_for_worker(vllm_config)
kv_publishers = []

for dp_rank in range(data_parallel_size):
for dp_rank in range(dp_start, dp_start + dp_size):
if consolidator_enabled:
# TODO: Use different port for each dp_rank once KVBM supports DP
zmq_endpoint = f"tcp://127.0.0.1:{consolidator_port}"
Expand Down Expand Up @@ -561,8 +549,9 @@ async def register_vllm_model(
runtime_config.reasoning_parser = config.dyn_reasoning_parser

# Get data_parallel_size from vllm_config (defaults to 1)
data_parallel_size = getattr(vllm_config.parallel_config, "data_parallel_size", 1)
runtime_config.data_parallel_size = data_parallel_size
dp_range = get_dp_range_for_worker(vllm_config)
runtime_config.data_parallel_start_rank = dp_range[0]
runtime_config.data_parallel_size = dp_range[1]

# Configure media decoder for frontend image decoding when enabled
# This enables frontend to decode images and transfer via NIXL RDMA
Expand Down Expand Up @@ -675,10 +664,6 @@ async def init_prefill(
runtime.register_engine_route("wake_up", handler.wake_up)
logger.info("Registered engine routes: /engine/sleep, /engine/wake_up")

# Handle non-leader nodes - don't serve endpoints
if config.engine_args.data_parallel_rank:
await _handle_non_leader_node(config.engine_args.data_parallel_rank)
return
shutdown_endpoints[:] = [generate_endpoint, clear_endpoint]

# Register prefill model with ModelType.Prefill
Expand Down Expand Up @@ -791,15 +776,13 @@ async def init(
factory = StatLoggerFactory(
endpoint=generate_endpoint,
component_gauges=component_gauges,
dp_rank=config.engine_args.data_parallel_rank or 0,
)
else:
# Factory is created without component_gauges; setup_vllm_engine() will
# create the gauges after setup_multiprocess_prometheus() and set them
# on the factory before vLLM calls create_stat_logger().
factory = StatLoggerFactory(
endpoint=generate_endpoint,
dp_rank=config.engine_args.data_parallel_rank or 0,
)
(
engine_client,
Expand Down Expand Up @@ -858,11 +841,6 @@ async def init(
runtime.register_engine_route("wake_up", handler.wake_up)
logger.info("Registered engine routes: /engine/sleep, /engine/wake_up")

# Handle non-leader nodes - don't serve endpoints
if config.engine_args.data_parallel_rank:
await _handle_non_leader_node(config.engine_args.data_parallel_rank)
return

# Parse endpoint types from --endpoint-types flag
model_type = parse_endpoint_types(config.endpoint_types)
logger.info(f"Registering model with endpoint types: {config.endpoint_types}")
Expand Down Expand Up @@ -1011,11 +989,6 @@ async def init_omni(
# Set up metrics collection for vLLM and LMCache metrics
setup_metrics_collection(config, generate_endpoint, logger)

# Handle non-leader nodes - don't serve endpoints
if config.engine_args.data_parallel_rank:
await _handle_non_leader_node(config.engine_args.data_parallel_rank)
return

# TODO: extend for multi-stage pipelines
model_type = get_output_modalities(config.output_modalities, config.model)
if model_type is None:
Expand Down
23 changes: 0 additions & 23 deletions components/src/dynamo/vllm/publisher.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,24 +19,6 @@
DYNAMO_COMPONENT_REGISTRY = CollectorRegistry()


class NullStatLogger(StatLoggerBase):
def __init__(self):
pass

def record(
self,
scheduler_stats: Optional[SchedulerStats],
iteration_stats: Optional[IterationStats],
engine_idx: int = 0,
*args,
**kwargs,
):
pass

def log_engine_initialized(self):
pass


class DynamoStatLoggerPublisher(StatLoggerBase):
"""Stat logger publisher. Wrapper for the WorkerMetricsPublisher to match the StatLoggerBase interface."""

Expand Down Expand Up @@ -106,22 +88,17 @@ def __init__(
self,
endpoint: Endpoint,
component_gauges: Optional[LLMBackendMetrics] = None,
dp_rank: int = 0,
) -> None:
self.endpoint = endpoint
self.component_gauges = component_gauges
self.created_logger: Optional[DynamoStatLoggerPublisher] = None
self.dp_rank = dp_rank

def create_stat_logger(self, dp_rank: int) -> StatLoggerBase:
if self.dp_rank != dp_rank:
return NullStatLogger()
# component_gauges must be set by setup_vllm_engine() before vLLM
# calls create_stat_logger() during engine initialization.
assert (
self.component_gauges is not None
), "component_gauges must be set before creating stat loggers"

logger = DynamoStatLoggerPublisher(
endpoint=self.endpoint,
dp_rank=dp_rank,
Expand Down
20 changes: 10 additions & 10 deletions examples/backends/vllm/launch/dep.sh
Original file line number Diff line number Diff line change
Expand Up @@ -35,16 +35,16 @@ python -m dynamo.frontend --router-mode kv &
# Routing to DP workers managed by Dynamo
# Chose Qwen3-30B because its a small MOE that can fit on smaller GPUs (L40S for example)
# --enforce-eager is added for quick deployment. for production use, need to remove this flag
for i in {0..3}; do
VLLM_NIXL_SIDE_CHANNEL_PORT=$((20096 + i)) \
CUDA_VISIBLE_DEVICES=$i python3 -m dynamo.vllm \
--model "$MODEL" \
--data-parallel-rank $i \
--data-parallel-size 4 \
--enable-expert-parallel \
--enforce-eager \
--kv-events-config "{\"publisher\":\"zmq\",\"topic\":\"kv-events\",\"endpoint\":\"tcp://*:$((20080 + i))\",\"enable_kv_cache_events\":true}" &
done
VLLM_NIXL_SIDE_CHANNEL_PORT=20096 \
python3 -m dynamo.vllm \
--model Qwen/Qwen3-30B-A3B \
--data-parallel-hybrid-lb \
--data-parallel-size 4 \
--data-parallel-size-local 4 \
--data-parallel-start-rank 0 \
--enable-expert-parallel \
--enforce-eager \
--kv-events-config "{\"publisher\":\"zmq\",\"topic\":\"kv-events\",\"endpoint\":\"tcp://*:20080\",\"enable_kv_cache_events\":true}" &

echo "All workers starting. (press Ctrl+C to stop)..."
wait
37 changes: 18 additions & 19 deletions examples/backends/vllm/launch/dsr1_dep.sh
Original file line number Diff line number Diff line change
Expand Up @@ -116,25 +116,24 @@ mkdir -p $LOG_DIR
# the GPU memory requires for vLLM reservation and runtime spike (not
# reserved by vLLM) can be different and cause model fails to start,
# adjust '--gpu-memory-utilization' as needed
for ((i=0; i<GPUS_PER_NODE; i++)); do
dp_rank=$((i + NODE_RANK * GPUS_PER_NODE))
CUDA_VISIBLE_DEVICES=$i \
VLLM_NIXL_SIDE_CHANNEL_PORT=$((20096 + i)) \
VLLM_ALL2ALL_BACKEND="deepep_low_latency" \
VLLM_USE_DEEP_GEMM=1 \
VLLM_RANDOMIZE_DP_DUMMY_INPUTS=1 \
python3 -m dynamo.vllm \
--model $MODEL \
--data_parallel_size $DATA_PARALLEL_SIZE \
--data-parallel-rank $dp_rank \
--enable-expert-parallel \
--max-model-len 4096 \
--data-parallel-address $MASTER_ADDR \
--data-parallel-rpc-port 13345 \
--gpu-memory-utilization 0.91 \
--enforce-eager \
--kv-events-config "{\"publisher\":\"zmq\",\"topic\":\"kv-events\",\"endpoint\":\"tcp://*:$((20080 + i))\",\"enable_kv_cache_events\":true}" 2>&1 | tee $LOG_DIR/dsr1_dep_${dp_rank}.log &
done
dp_start_rank=$((NODE_RANK * GPUS_PER_NODE))
VLLM_NIXL_SIDE_CHANNEL_PORT=20096 \
VLLM_ALL2ALL_BACKEND="deepep_low_latency" \
VLLM_USE_DEEP_GEMM=1 \
VLLM_RANDOMIZE_DP_DUMMY_INPUTS=1 \
python3 -m dynamo.vllm \
--model $MODEL \
--data-parallel-hybrid-lb \
--data-parallel-size $DATA_PARALLEL_SIZE \
--data-parallel-size-local $GPUS_PER_NODE \
--data-parallel-start-rank $dp_start_rank \
--enable-expert-parallel \
--max-model-len 4096 \
--data-parallel-address $MASTER_ADDR \
--data-parallel-rpc-port 13345 \
--gpu-memory-utilization 0.91 \
--enforce-eager \
--kv-events-config "{\"publisher\":\"zmq\",\"topic\":\"kv-events\",\"endpoint\":\"tcp://*:20080\",\"enable_kv_cache_events\":true}" 2>&1 | tee $LOG_DIR/dsr1_dep_${dp_start_rank}.log &

echo "All workers starting. (press Ctrl+C to stop)..."
wait
5 changes: 5 additions & 0 deletions lib/bindings/python/rust/llm/local_model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,11 @@ impl ModelRuntimeConfig {
self.inner.reasoning_parser = reasoning_parser;
}

#[setter]
fn set_data_parallel_start_rank(&mut self, data_parallel_start_rank: u32) {
self.inner.data_parallel_start_rank = data_parallel_start_rank;
}

#[setter]
fn set_data_parallel_size(&mut self, data_parallel_size: u32) {
self.inner.data_parallel_size = data_parallel_size;
Expand Down
5 changes: 3 additions & 2 deletions lib/kv-router/benches/active_sequences_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -289,11 +289,12 @@ async fn run_benchmark(
// Total bench workers = trace workers × duplication factor.
// Each gets a unique WorkerWithDpRank in the shared multi-worker.
let total_workers = num_trace_workers * inference_worker_duplication_factor;
let dp_sizes: HashMap<u64, u32> = (0..total_workers as u64).map(|id| (id, 1)).collect();
let dp_range: HashMap<u64, (u32, u32)> =
(0..total_workers as u64).map(|id| (id, (0, 1))).collect();
let multi = Arc::new(ActiveSequencesMultiWorker::new(
NoopSequencePublisher,
block_size as usize,
dp_sizes,
dp_range,
false,
0,
"bench",
Expand Down
1 change: 1 addition & 0 deletions lib/kv-router/src/protocols.rs
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ pub fn compute_seq_hash_for_block(block_hashes: &[LocalBlockHash]) -> Vec<Sequen
///
/// `ModelRuntimeConfig` (in `lib/llm`) implements this directly so no adapter type is needed.
pub trait WorkerConfigLike {
fn data_parallel_start_rank(&self) -> u32;
fn data_parallel_size(&self) -> u32;
fn max_num_batched_tokens(&self) -> Option<u64>;
fn total_kv_blocks(&self) -> Option<u64>;
Expand Down
8 changes: 5 additions & 3 deletions lib/kv-router/src/scheduling/queue.rs
Original file line number Diff line number Diff line change
Expand Up @@ -209,11 +209,12 @@ impl<P: SequencePublisher + 'static, C: WorkerConfigLike> SchedulerQueue<P, C> {

for (&worker_id, config) in configs.iter() {
let dp_size = config.data_parallel_size();
let dp_start_rank = config.data_parallel_start_rank();
let max_batched = config
.max_num_batched_tokens()
.unwrap_or(DEFAULT_MAX_BATCHED_TOKENS);

for dp_rank in 0..dp_size {
for dp_rank in dp_start_rank..dp_start_rank + dp_size {
let worker = WorkerWithDpRank::new(worker_id, dp_rank);
let tokens = active_tokens.get(&worker).copied().unwrap_or(0);
if (tokens as f64) <= threshold * (max_batched as f64) {
Expand Down Expand Up @@ -247,11 +248,12 @@ mod tests {
Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
) {
let dp_sizes: HashMap<u64, u32> = (0..num_workers as u64).map(|id| (id, 1)).collect();
let dp_range: HashMap<u64, (u32, u32)> =
(0..num_workers as u64).map(|id| (id, (0, 1))).collect();
let slots = Arc::new(ActiveSequencesMultiWorker::new(
NoopSequencePublisher,
block_size as usize,
dp_sizes,
dp_range,
false,
0,
"test",
Expand Down
Loading
Loading