Skip to content
Merged
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
14 changes: 11 additions & 3 deletions vime/backends/vllm_utils/vllm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -455,6 +455,11 @@ def build_vllm_cmd_and_env(server_args: dict[str, Any]) -> tuple[list[str], dict
"vime.backends.megatron_utils.update_weight.update_weight_from_tensor.vLLMColocateWorkerExtension",
]

worker_type = server_args.get("worker_type", "regular")
if worker_type in ("prefill", "decode") and topology.node_rank == 0:
env["VLLM_NIXL_SIDE_CHANNEL_HOST"] = host_for_subprocess
env["VLLM_NIXL_SIDE_CHANNEL_PORT"] = str(server_args["disaggregation_bootstrap_port"])

_forward_vllm_cli_args(args, cmd)
logger.info("Launching vLLM server: %s", redact_cmd_for_log(cmd))
return cmd, env
Expand Down Expand Up @@ -548,9 +553,8 @@ def init(
router_ip=None,
router_port=None,
):
# ``nccl_port`` / ``disaggregation_bootstrap_port`` are allocated by rollout
# port allocation but not consumed by vLLM (rendezvous uses ``dist_init_addr``).
del nccl_port, disaggregation_bootstrap_port
# ``nccl_port`` is allocated by rollout but unused by vLLM (rendezvous uses ``dist_init_addr``).
del nccl_port

gpus_per_engine = self.num_gpus_per_engine or self.args.rollout_num_gpus_per_engine
host = host or get_host_info()[1]
Expand All @@ -565,6 +569,7 @@ def init(
base_gpu_id=self.base_gpu_id,
vllm_overrides=self.vllm_overrides,
num_gpus_per_engine=gpus_per_engine,
disaggregation_bootstrap_port=disaggregation_bootstrap_port,
)
self._topology = self._server_args["topology"]
self.node_rank = self._topology.node_rank
Expand All @@ -575,6 +580,7 @@ def init(
self.router_port = router_port
self.server_host = self._server_args["host"]
self.server_port = port
self.disaggregation_bootstrap_port = disaggregation_bootstrap_port

if self.worker_type != "regular":
logger.warning(
Expand Down Expand Up @@ -1101,6 +1107,7 @@ def _compute_server_args(
base_gpu_id: int | None = None,
vllm_overrides: dict | None = None,
num_gpus_per_engine: int | None = None,
disaggregation_bootstrap_port: int | None = None,
) -> dict[str, Any]:
"""Build per-actor launch config for ``launch_server_process``."""
gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine
Expand Down Expand Up @@ -1140,6 +1147,7 @@ def _compute_server_args(
"pp_size": topology.pipeline_parallel_size,
"dp_size": _get_vllm_dp_size(args),
"seed": getattr(args, "seed", 1234) + rank,
"disaggregation_bootstrap_port": disaggregation_bootstrap_port,
}
_apply_vllm_overrides(args, server_args, vllm_overrides, rank)
return server_args
89 changes: 63 additions & 26 deletions vime/ray/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -884,7 +884,7 @@ def addr():
addr_and_ports[current_rank]["port"] = get_port()
addr_and_ports[current_rank]["nccl_port"] = get_port()

if worker_type == "prefill":
if worker_type in ("prefill", "decode"):
addr_and_ports[current_rank]["disaggregation_bootstrap_port"] = get_port()

if _gpus_per_engine > args.num_gpus_per_node:
Expand All @@ -907,25 +907,32 @@ def addr():
return addr_and_ports, node_port_cursor


def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool = False) -> tuple[str, int]:
def _start_router(
args,
*,
has_pd_disaggregation: bool = False,
force_new: bool = False,
bind: tuple[str, int] | None = None,
prefill_urls: list | None = None,
decode_urls: list | None = None,
) -> tuple[str, int, int]:
"""Start the rollout HTTP gateway (vllm-router)."""
if not force_new and args.vllm_router_ip is not None:
return args.vllm_router_ip, args.vllm_router_port

router_ip = _wrap_ipv6(get_host_info()[1])
if force_new:
router_port = find_available_port(random.randint(3000, 4000))
if bind is not None:
router_ip, router_port = bind
else:
router_port = args.vllm_router_port
if router_port is None:
if not force_new and args.vllm_router_ip is not None:
return args.vllm_router_ip, args.vllm_router_port, None
router_ip = _wrap_ipv6(get_host_info()[1])
if force_new or args.vllm_router_port is None:
router_port = find_available_port(random.randint(3000, 4000))
else:
router_port = args.vllm_router_port

from vllm_router.router_args import RouterArgs

from vime.utils.http_utils import run_router

router_args = RouterArgs.from_cli_args(args, use_router_prefix=True)

router_args.host = router_ip
router_args.port = router_port
router_args.prometheus_port = find_available_port(random.randint(4000, 5000))
Expand All @@ -934,20 +941,19 @@ def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool

if has_pd_disaggregation:
router_args.vllm_pd_disaggregation = True
# Disable circuit breaker to prevent RDMA transfer timeouts from
# marking decode workers as dead. Timeouts are transient (PCIe
# contention under high load) and do not indicate a dead server.
# Disable circuit breaker so transient RDMA transfer timeouts (PCIe
# contention under load) don't mark decode workers dead.
router_args.disable_circuit_breaker = True

if prefill_urls is not None:
router_args.prefill_urls = prefill_urls
router_args.decode_urls = decode_urls

logger.info(f"Launch router with args: {router_args}")

process = multiprocessing.Process(
target=run_router,
args=(router_args,),
)
process.daemon = True # Set the process as a daemon
process = multiprocessing.Process(target=run_router, args=(router_args,))
process.daemon = True
process.start()
# Wait 3 seconds
time.sleep(3)
assert process.is_alive()
logger.info(f"Router launched at {router_ip}:{router_port}, Prometheus port: {router_args.prometheus_port}")
Expand Down Expand Up @@ -997,9 +1003,17 @@ def start_rollout_servers(args, pg) -> dict[str, RolloutServer]:
model_cfg.resolve(args)

has_pd = model_cfg.has_pd_disaggregation
router_ip, router_port, prom_port = _start_router(
args, has_pd_disaggregation=has_pd, force_new=(model_idx > 0)
)
use_static_pd_router = has_pd
if use_static_pd_router:
router_ip = _wrap_ipv6(get_host_info()[1])
router_port = find_available_port(random.randint(3000, 4000))
prom_port = None # assigned when the router actually launches, after URL collection
engine_router_ip, engine_router_port = None, None
else:
router_ip, router_port, prom_port = _start_router(
args, has_pd_disaggregation=has_pd, force_new=(model_idx > 0)
)
engine_router_ip, engine_router_port = router_ip, router_port

# Write back so downstream readers (vllm_rollout, vllm_engine) see the
# router we just started (only relevant for first model in multi-model setups).
Expand Down Expand Up @@ -1056,7 +1070,7 @@ def _make_group(group_cfg, router_ip, router_port, overrides_extra=None):
for group_cfg in model_cfg.server_groups:
if group_cfg.worker_type != "encoder":
continue
group = _make_group(group_cfg, router_ip, router_port)
group = _make_group(group_cfg, engine_router_ip, engine_router_port)
handles, port_cursors = group.start_engines(port_cursors)
if handles:
ray.get(handles)
Expand All @@ -1077,7 +1091,7 @@ def _make_group(group_cfg, router_ip, router_port, overrides_extra=None):
if encoder_urls and group_cfg.worker_type in ("prefill", "regular"):
overrides_extra["language_only"] = True
overrides_extra["encoder_urls"] = encoder_urls
group = _make_group(group_cfg, router_ip, router_port, overrides_extra=overrides_extra)
group = _make_group(group_cfg, engine_router_ip, engine_router_port, overrides_extra=overrides_extra)
handles, port_cursors = group.start_engines(port_cursors)
non_encoder_handles.extend(handles)
server_groups.append(group)
Expand All @@ -1088,14 +1102,37 @@ def _make_group(group_cfg, router_ip, router_port, overrides_extra=None):
# No EPD — start all groups in one pass (original path).
all_init_handles: list = []
for group_cfg in model_cfg.server_groups:
group = _make_group(group_cfg, router_ip, router_port)
group = _make_group(group_cfg, engine_router_ip, engine_router_port)
handles, port_cursors = group.start_engines(port_cursors)
all_init_handles.extend(handles)
server_groups.append(group)

if all_init_handles:
ray.get(all_init_handles)

if use_static_pd_router:
prefill_urls: list[tuple] = []
decode_urls: list[str] = []
for g in server_groups:
for e in g.engines:
if e is None:
continue
if g.worker_type == "prefill":
url = ray.get(e.get_url.remote())
if url:
prefill_urls.append((url, None))
elif g.worker_type == "decode":
url = ray.get(e.get_url.remote())
if url:
decode_urls.append(url)
_, _, prom_port = _start_router(
args,
has_pd_disaggregation=True,
bind=(router_ip, router_port),
prefill_urls=prefill_urls,
decode_urls=decode_urls,
)

servers[model_cfg.name] = RolloutServer(
server_groups=server_groups,
router_ip=router_ip,
Expand Down
Loading