[EPD] Support dynamic encoder register (#22253)

This commit is contained in:
siyu
2026-06-08 16:05:16 +08:00
committed by GitHub
parent 18d728967a
commit 13dda3b8de
6 changed files with 442 additions and 19 deletions
@@ -9,6 +9,7 @@ import weakref
from abc import ABC, abstractmethod
from array import array
from collections import OrderedDict, defaultdict
from contextlib import asynccontextmanager
from enum import IntEnum
from http import HTTPStatus
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
@@ -16,8 +17,12 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import aiohttp
import numpy as np
import torch
import uvicorn
import zmq
import zmq.asyncio
from aiohttp import ClientSession, ClientTimeout
from fastapi import FastAPI
from fastapi.responses import ORJSONResponse, Response
from transformers import PretrainedConfig
from sglang.srt.distributed.parallel_state import (
@@ -44,6 +49,219 @@ if TYPE_CHECKING:
from sglang.srt.managers.scheduler import Scheduler
class EncoderBootstrapServer:
"""Lightweight bootstrap server for dynamic encoder discovery.
Built on FastAPI + uvicorn to match the style of
:mod:`sglang.srt.entrypoints.http_server`. Runs in a daemon thread so
the language-only tokenizer manager's main loop is unblocked.
The set of registered URLs is exposed as the ``urls`` list passed in at
construction time. Callers that want to observe registrations without
going through HTTP -- typically a co-located :class:`MMReceiver` -- share
that list by reference: register/unregister mutate it in place under an
internal lock, and the receiver simply reads ``self.encode_urls`` (the
same list). When ``urls`` is ``None`` the server allocates its own list,
accessible through :meth:`list_urls`.
Health-check tuning is controlled by env vars
``SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_INTERVAL`` (seconds; 0 disables)
and ``SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_TIMEOUT`` (seconds). Explicit
constructor args take precedence over the env vars.
"""
def __init__(
self,
host: str,
port: int,
urls: Optional[List[str]] = None,
health_check_interval: Optional[float] = None,
health_check_timeout: Optional[float] = None,
):
self.host = host
self.port = port
self._urls: List[str] = urls if urls is not None else []
self._lock = threading.Lock()
self._server: Optional["uvicorn.Server"] = None # set in _run_server
self._health_check_interval = (
health_check_interval
if health_check_interval is not None
else envs.SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_INTERVAL.get()
)
self._health_check_timeout = (
health_check_timeout
if health_check_timeout is not None
else envs.SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_TIMEOUT.get()
)
self._consecutive_failures: Dict[str, int] = {}
self._max_consecutive_failures = 3
@asynccontextmanager
async def lifespan(fast_api_app: FastAPI):
task: Optional[asyncio.Task] = None
if self._health_check_interval > 0:
task = asyncio.create_task(self._health_check_loop())
try:
yield
finally:
if task is not None:
task.cancel()
try:
await task
except (asyncio.CancelledError, Exception):
pass
self.app = FastAPI(lifespan=lifespan, openapi_url=None)
@self.app.get("/health")
async def _health() -> Response:
return Response("OK")
@self.app.post("/register_encoder_url")
async def _register(data: dict):
url = data.get("url") if isinstance(data, dict) else None
if not url:
return ORJSONResponse(
{"error": "Missing or empty 'url' field"}, status_code=400
)
self.register(url)
return Response("OK")
@self.app.delete("/unregister_encoder_url")
async def _unregister(data: dict):
url = data.get("url") if isinstance(data, dict) else None
if not url:
return ORJSONResponse(
{"error": "Missing or empty 'url' field"}, status_code=400
)
self.unregister(url)
return Response("OK")
@self.app.get("/list_encoder_urls")
async def _list():
return {"encoder_urls": self.list_urls()}
self.thread = threading.Thread(
target=self._run_server, daemon=True, name="EncoderBootstrap"
)
self.thread.start()
# ------------------------------------------------------------------ #
# In-process API (thread-safe; safe to call from any thread) #
# ------------------------------------------------------------------ #
def register(self, url: str) -> bool:
"""Add *url* if not already present. Returns True if added."""
with self._lock:
if url not in self._urls:
self._urls.append(url)
self._consecutive_failures.pop(url, None)
logger.info(f"Registered encoder URL: {url}")
return True
logger.debug(f"Encoder URL already registered: {url}")
return False
def unregister(self, url: str) -> bool:
"""Remove *url* if present. Returns True if removed."""
with self._lock:
if url in self._urls:
self._urls.remove(url)
logger.info(f"Unregistered encoder URL: {url}")
return True
return False
def list_urls(self) -> List[str]:
"""Return a snapshot of all registered encoder URLs."""
with self._lock:
return list(self._urls)
# ------------------------------------------------------------------ #
# Health check #
# ------------------------------------------------------------------ #
async def _probe(self, session, url: str) -> bool:
try:
async with session.get(f"{url}/health") as resp:
return resp.status == 200
except Exception:
return False
async def _health_check_loop(self):
"""Probe each registered encoder periodically and evict dead ones."""
timeout = ClientTimeout(total=self._health_check_timeout)
while True:
try:
await asyncio.sleep(self._health_check_interval)
snapshot = self.list_urls()
if not snapshot:
continue
async with ClientSession(timeout=timeout) as session:
results = await asyncio.gather(
*(self._probe(session, url) for url in snapshot),
return_exceptions=True,
)
evicted = []
with self._lock:
for url, ok in zip(snapshot, results):
if ok is True:
self._consecutive_failures.pop(url, None)
else:
self._consecutive_failures[url] = (
self._consecutive_failures.get(url, 0) + 1
)
if (
self._consecutive_failures[url]
>= self._max_consecutive_failures
):
if url in self._urls:
self._urls.remove(url)
self._consecutive_failures.pop(url, None)
evicted.append(url)
if evicted:
logger.warning(
f"Health check evicted {len(evicted)} encoder(s) "
f"after {self._max_consecutive_failures} consecutive "
f"failures: {evicted}"
)
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(f"Health check loop error: {e}", exc_info=True)
# ------------------------------------------------------------------ #
# Lifecycle #
# ------------------------------------------------------------------ #
def _run_server(self):
config = uvicorn.Config(
self.app,
host=self.host,
port=self.port,
log_level="warning",
access_log=False,
loop="auto",
)
self._server = uvicorn.Server(config)
logger.info(
f"EncoderBootstrapServer starting on {self.host}:{self.port} "
f"(health_check every {self._health_check_interval}s, "
f"timeout {self._health_check_timeout}s)"
)
try:
self._server.run()
except Exception as e:
logger.error(f"EncoderBootstrapServer error: {e}", exc_info=True)
def close(self):
if self._server is not None:
# uvicorn polls should_exit on its own event loop; thread-safe.
self._server.should_exit = True
logger.info("Stopping EncoderBootstrapServer...")
if self.thread.is_alive():
self.thread.join(timeout=5)
logger.info("EncoderBootstrapServer thread stopped")
def _grpc_target(url: str) -> str:
if url.startswith("grpc://"):
return url[len("grpc://") :]
@@ -1198,10 +1416,18 @@ class MMReceiverBase(ABC):
tp_rank: Optional[int] = None,
tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
encode_urls: Optional[List[str]] = None,
):
self.context = zmq.asyncio.Context(20)
self.encoder_transfer_backend = server_args.encoder_transfer_backend
self.encode_urls = server_args.encoder_urls
# When ``encode_urls`` is shared with an :class:`EncoderBootstrapServer`
# (tokenizer manager process), it grows / shrinks in place as encoders
# register or unregister; the receiver always sees the current set.
# When None (e.g. in a scheduler subprocess that has no in-process
# bootstrap), fall back to a snapshot of the static --encoder-urls.
self.encode_urls: List[str] = (
encode_urls if encode_urls is not None else list(server_args.encoder_urls)
)
self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
self.host = get_local_ip_auto(server_args.host)
self.pp_rank = pp_rank
@@ -1328,7 +1554,13 @@ class MMReceiverBase(ABC):
):
req_id = None
try:
if len(self.encode_urls) == 0 or not need_wait_for_mm_inputs:
# ``self.encode_urls`` is shared by reference with the bootstrap
# server (when running) so it always reflects the current set.
# Snapshot once for the duration of this request to avoid races
# against concurrent register / unregister.
encode_urls = list(self.encode_urls)
if len(encode_urls) == 0 or not need_wait_for_mm_inputs:
return None
req_id = uuid.uuid4().hex
embedding_port, recv_socket = get_zmq_socket_on_host(
@@ -1342,7 +1574,14 @@ class MMReceiverBase(ABC):
)
send_time = time.monotonic()
asyncio.create_task(
self.encode(req_id, mm_data, embedding_port, "encode", "send")
self.encode(
req_id,
mm_data,
embedding_port,
"encode",
"send",
encode_urls=encode_urls,
)
)
result = await asyncio.wait_for(
self._recv_mm_data(req_id, recv_socket, mm_processor, prompt),
@@ -1451,14 +1690,26 @@ class MMReceiverBase(ABC):
mm_data = self._extract_url_data(obj)
if obj.rid is None:
obj.rid = uuid.uuid4().hex
if mm_data and self.encode_urls:
logger.info(f"Processing {len(mm_data)} mm items for request {obj.rid}")
# ``self.encode_urls`` is the shared list maintained by the bootstrap
# server (and pre-populated with --encoder-urls); take a snapshot for
# the duration of this dispatch.
encode_urls = list(self.encode_urls)
if mm_data and encode_urls:
logger.info(
f"Dispatching {len(mm_data)} mm items to {len(encode_urls)} "
f"encoder(s) {encode_urls} for request {obj.rid}"
)
obj.need_wait_for_mm_inputs = True
num_items_assigned = self._assign_items_by_modality(
mm_data, len(self.encode_urls)
mm_data, len(encode_urls)
)
obj.num_items_assigned = num_items_assigned
# Freeze the encoder URL snapshot onto obj so the scheduler
# subprocess uses the same list when indexing encoder_idx.
obj.encoder_urls = encode_urls
# For mooncake, No tokenizer-side thread.
# Save mm_data (extracted URL list) onto obj so the scheduler-side
@@ -1477,10 +1728,22 @@ class MMReceiverBase(ABC):
"encode",
num_items_assigned,
None,
encode_urls,
),
daemon=True,
)
encode_thread.start()
else:
# No encoder URLs available (bootstrap may not have any registered yet);
# reset the flag so the scheduler does not wait for embeddings that will
# never arrive. A warning is emitted so the user can diagnose why
# disaggregation is not happening for this request.
if mm_data:
logger.warning(
f"No encoder URLs available for request {obj.rid}; "
"processing without encoder disaggregation."
)
obj.need_wait_for_mm_inputs = False
# For zmq_to_scheduler
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
@@ -1490,11 +1753,18 @@ class MMReceiverBase(ABC):
isinstance(recv_req, TokenizedGenerateReqInput)
and recv_req.need_wait_for_mm_inputs is True
):
# Use the URL snapshot frozen by the tokenizer when it
# computed num_items_assigned -- the encoder_idx values in
# that assignment must index into this exact list. Falling
# back to ``self.encode_urls`` would only matter if the
# tokenizer never set encoder_urls (legacy / static path).
encode_urls = recv_req.encoder_urls or list(self.encode_urls)
waiting_req = waiting_cls(
rid=recv_req.rid,
recv_req=recv_req,
mm_processor=self.mm_processor,
encoder_urls=self.encode_urls,
encoder_urls=encode_urls,
model_type=self.model_type,
host_name=self.hostname,
receive_count=self.tp_size,
@@ -1561,7 +1831,13 @@ class MMReceiverBase(ABC):
return new_recv_reqs, abort_reqs
def _run_encode_in_thread(
self, req_id, mm_data, endpoint_encode, num_items_assigned, embedding_port
self,
req_id,
mm_data,
endpoint_encode,
num_items_assigned,
embedding_port,
encode_urls=None,
):
try:
asyncio.run(
@@ -1572,6 +1848,7 @@ class MMReceiverBase(ABC):
endpoint_encode=endpoint_encode,
endpoint_send=None,
num_items_assigned=num_items_assigned,
encode_urls=encode_urls,
)
)
except Exception as e:
@@ -1724,6 +2001,7 @@ class MMReceiverHTTP(MMReceiverBase):
tp_rank: Optional[int] = None,
tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
encode_urls: Optional[List[str]] = None,
):
super().__init__(
server_args,
@@ -1733,6 +2011,7 @@ class MMReceiverHTTP(MMReceiverBase):
tp_rank=tp_rank,
tp_group=tp_group,
scheduler=scheduler,
encode_urls=encode_urls,
)
# For zmq_to_scheduler and mooncake
@@ -1754,10 +2033,13 @@ class MMReceiverHTTP(MMReceiverBase):
for i, response in enumerate(responses):
if isinstance(response, asyncio.TimeoutError):
timeout_val = envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()
encoder_label = encode_requests[i].get(
"encoder_url", f"idx={encode_requests[i].get('encoder_idx')}"
)
logger.error(
f"Encoder HTTP request timeout ({timeout_val}s) for req_id={req_id} "
f"(request {i}), "
f"encoder={self.encode_urls[encode_requests[i]['encoder_idx']]}"
f"encoder={encoder_label}"
)
return False
elif isinstance(response, Exception):
@@ -1785,10 +2067,13 @@ class MMReceiverHTTP(MMReceiverBase):
endpoint_encode,
endpoint_send,
num_items_assigned=None,
encode_urls=None,
):
if len(mm_data) == 0:
return
effective_urls = encode_urls if encode_urls is not None else self.encode_urls
# get unique modalities with order preserved
modalities = [mm_item.get("modality") for mm_item in mm_data]
modalities = list(dict.fromkeys(modalities))
@@ -1796,7 +2081,7 @@ class MMReceiverHTTP(MMReceiverBase):
if num_items_assigned is None:
num_items_assigned = self._assign_items_by_modality(
mm_data, len(self.encode_urls)
mm_data, len(effective_urls)
)
# Calculate total num_parts across all modalities
@@ -1822,6 +2107,7 @@ class MMReceiverHTTP(MMReceiverBase):
encode_requests.append(
{
"encoder_idx": idx,
"encoder_url": effective_urls[idx],
"mm_items": [
mm_item.get("url")
for mm_item in mm_data_modality[
@@ -1847,7 +2133,7 @@ class MMReceiverHTTP(MMReceiverBase):
tasks = [
session.post(
f"{self.encode_urls[encode_request['encoder_idx']]}/{endpoint_encode}",
f"{effective_urls[encode_request['encoder_idx']]}/{endpoint_encode}",
json=encode_request,
)
for encode_request in encode_requests
@@ -1891,7 +2177,7 @@ class MMReceiverHTTP(MMReceiverBase):
)
metadata_tasks.append(
session.post(
f"{self.encode_urls[response_json['encoder_idx']]}/{endpoint_send}",
f"{effective_urls[response_json['encoder_idx']]}/{endpoint_send}",
json=response_json,
)
)
@@ -1909,6 +2195,7 @@ class MMReceiverGrpc(MMReceiverBase):
tp_rank: Optional[int] = None,
tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
encode_urls: Optional[List[str]] = None,
):
super().__init__(
server_args,
@@ -1918,6 +2205,7 @@ class MMReceiverGrpc(MMReceiverBase):
tp_rank=tp_rank,
tp_group=tp_group,
scheduler=scheduler,
encode_urls=encode_urls,
)
def build_and_send_encode_request(self, image_urls, rid):
@@ -1940,10 +2228,13 @@ class MMReceiverGrpc(MMReceiverBase):
endpoint_encode,
endpoint_send,
num_items_assigned=None,
encode_urls=None,
):
if not mm_data:
return
effective_urls = encode_urls if encode_urls is not None else self.encode_urls
# gRPC currently only supports image; flatten new dict formats to simple lists
if mm_data and isinstance(mm_data[0], dict):
non_image = [
@@ -1963,10 +2254,10 @@ class MMReceiverGrpc(MMReceiverBase):
encode_requests = []
if num_items_assigned is None:
encode_idx = list(range(len(self.encode_urls)))
encode_idx = list(range(len(effective_urls)))
random.shuffle(encode_idx)
num_items_assigned = [
(idx + len(img_data)) // len(self.encode_urls) for idx in encode_idx
(idx + len(img_data)) // len(effective_urls) for idx in encode_idx
]
num_parts = sum(1 for x in num_items_assigned if x != 0)
cum_num_items = 0
@@ -1993,7 +2284,7 @@ class MMReceiverGrpc(MMReceiverBase):
grpc_tasks = [
asyncio.to_thread(
_grpc_encode_request,
_grpc_target(self.encode_urls[encode_request["encoder_idx"]]),
_grpc_target(effective_urls[encode_request["encoder_idx"]]),
encode_request,
)
for encode_request in encode_requests
@@ -2039,7 +2330,7 @@ class MMReceiverGrpc(MMReceiverBase):
grpc_metadata_tasks.append(
asyncio.to_thread(
_grpc_send_request,
_grpc_target(self.encode_urls[response_json["encoder_idx"]]),
_grpc_target(effective_urls[response_json["encoder_idx"]]),
response_json,
)
)
@@ -2084,12 +2375,13 @@ def create_mm_receiver(
tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
transport_mode: Optional[str] = None,
encode_urls: Optional[List[str]] = None,
):
if transport_mode is None:
transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get()
logger.debug(f"MMReceiver transport_mode from env: {transport_mode}")
_validate_transport_mode(transport_mode, server_args.encoder_urls)
_validate_transport_mode(transport_mode, encode_urls or server_args.encoder_urls)
logger.info(f"EPD MMReceiver: using transport_mode={transport_mode}")
receiver_cls = _MM_RECEIVER_BY_MODE.get(transport_mode)
@@ -2103,4 +2395,5 @@ def create_mm_receiver(
tp_rank=tp_rank,
tp_group=tp_group,
scheduler=scheduler,
encode_urls=encode_urls,
)
@@ -8,6 +8,7 @@ import logging
import multiprocessing as mp
import os
import pickle
import threading
import time
import traceback
from collections import defaultdict
@@ -16,6 +17,7 @@ from typing import Dict, List, Optional, Set, Tuple, Union
import aiohttp
import numpy as np
import requests as http_requests
import torch
import uvicorn
import zmq
@@ -3127,6 +3129,74 @@ def launch_encoder(server_args, schedule_path, dist_init_method, rank):
traceback.print_exc()
def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
"""Asynchronously register this encoder with each bootstrap URL.
Spawns a daemon thread that retries each URL independently with bounded
backoff. The encoder's own startup is not blocked: if some bootstrap
server is slow or unreachable, only the background worker waits.
Inspired by ``_ensure_prefill_info`` in disaggregation/decode.py: each
target keeps its own retry count and is retried at a fixed interval
instead of serialising sleeps in a single thread.
"""
encoder_url = server_args.url()
payload = {"url": encoder_url}
bootstrap_urls = list(server_args.encoder_register_urls)
if not bootstrap_urls:
return
max_retries = 30
retry_interval = 5.0
request_timeout = 5.0
def _try_register_once(bootstrap_url: str) -> bool:
try:
resp = http_requests.post(
f"{bootstrap_url}/register_encoder_url",
json=payload,
timeout=request_timeout,
)
if resp.status_code == 200:
logger.info(
f"Registered encoder URL '{encoder_url}' with bootstrap "
f"at {bootstrap_url}"
)
return True
logger.warning(
f"Bootstrap {bootstrap_url} returned {resp.status_code}: {resp.text}"
)
except Exception as e:
logger.debug(f"Register attempt to {bootstrap_url} failed: {e}")
return False
def _worker():
pending = list(bootstrap_urls)
retry_count = {url: 0 for url in pending}
while pending:
still_pending = []
for bootstrap_url in pending:
if _try_register_once(bootstrap_url):
continue
retry_count[bootstrap_url] += 1
if retry_count[bootstrap_url] >= max_retries:
logger.error(
f"Giving up on bootstrap {bootstrap_url} after "
f"{max_retries} attempts. Encoder discovery via this "
f"bootstrap will be incomplete."
)
continue
still_pending.append(bootstrap_url)
pending = still_pending
if pending:
time.sleep(retry_interval)
threading.Thread(
target=_worker, daemon=True, name="encoder-bootstrap-register"
).start()
def launch_server(server_args: ServerArgs):
configure_logger(server_args, prefix=" encode_server")
if server_args.dp_size > 1:
@@ -3156,6 +3226,11 @@ def launch_server(server_args: ServerArgs):
daemon=True,
).start()
encoder = MMEncoder(server_args, dist_init_method=dist_init_method)
# Register this encoder's URL with prefill server(s) if configured.
if server_args.encoder_register_urls:
_register_encoder_url_with_bootstrap(server_args)
uvicorn.run(app, host=server_args.host, port=server_args.port)
@@ -3228,6 +3303,10 @@ def _launch_server_dp(server_args: ServerArgs):
worker_processes,
)
# Register this encoder's URL with prefill server(s) if configured.
if server_args.encoder_register_urls:
_register_encoder_url_with_bootstrap(server_args)
uvicorn.run(app, host=server_args.host, port=server_args.port)
+3
View File
@@ -840,6 +840,9 @@ class Envs:
SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU = EnvBool(False)
SGLANG_ENCODER_MAX_BATCH_SIZE = EnvInt(8)
SGLANG_ENCODER_PREPROC_WORKERS = EnvInt(8)
# EncoderBootstrapServer health-check tuning. Interval == 0 disables it.
SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_INTERVAL = EnvFloat(10.0)
SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_TIMEOUT = EnvFloat(2.0)
# Persistent receiver-side GPU embedding pool size for mooncake EPD transport.
# 0 disables (per-request register/deregister). 4096 = 4GB default per TP
SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096)
+7
View File
@@ -270,6 +270,9 @@ class GenerateReqInput(BaseReq):
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Snapshot of encoder URLs at the time tokenizer-side computed
# ``num_items_assigned``.
encoder_urls: Optional[List[str]] = None
# Multimodal tiling controls (extensions)
max_dynamic_patch: Optional[int] = None
@@ -815,6 +818,10 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Encoder URL snapshot frozen at tokenizer-side dispatch time so that
# encoder_idx assignments stay consistent in the scheduler subprocess.
# Internal IPC only.
encoder_urls: Optional[List[str]] = None
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
@@ -487,11 +487,27 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.fake_bootstrap_room_counter = 0
# Encoder Disaggregation
self.encoder_bootstrap_server = None
if self.server_args.language_only:
from sglang.srt.disaggregation.encode_receiver import (
EncoderBootstrapServer,
)
# Shared mutable URL list: the bootstrap server appends / removes
# entries as encoders register, the receiver reads from the same
# list. Pre-populated with static --encoder-urls so the legacy
# CLI flag still works (alongside dynamic registrations).
self.encoder_urls: List[str] = list(self.server_args.encoder_urls)
self.encoder_bootstrap_server = EncoderBootstrapServer(
host=self.server_args.host,
port=self.server_args.encoder_bootstrap_port,
urls=self.encoder_urls,
)
self.mm_receiver = create_mm_receiver(
self.server_args,
dtype=self.model_config.dtype,
hf_config=self.model_config.hf_config,
encode_urls=self.encoder_urls,
)
def init_metric_collector_watchdog(self):
@@ -1122,6 +1138,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
mm_data_mooncake=obj.mm_data_mooncake,
encoder_urls=obj.encoder_urls,
)
elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available
+26 -2
View File
@@ -842,6 +842,10 @@ class ServerArgs:
language_only: bool = False
encoder_transfer_backend: str = ENCODER_TRANSFER_BACKEND_CHOICES[0]
encoder_urls: List[str] = dataclasses.field(default_factory=list)
# Port of the standalone EncoderBootstrapServer started by the language-only
# tokenizer manager. Encoder workers register here.
encoder_bootstrap_port: int = 8997
encoder_register_urls: List[str] = dataclasses.field(default_factory=list)
enable_adaptive_dispatch_to_encoder: bool = False
# For model weight update and weight loading
@@ -3960,8 +3964,10 @@ class ServerArgs:
)
if self.language_only and len(self.encoder_urls) == 0:
raise ValueError(
"requires at least one encoder urls to be set via --encoder-urls"
logger.info(
"--language-only is set without --encoder-urls. Encoders are "
"expected to register dynamically via the "
"EncoderBootstrapServer."
)
# Validate IB devices when mooncake backend is used
@@ -7036,6 +7042,24 @@ class ServerArgs:
default=[],
help="List of encoder server urls.",
)
parser.add_argument(
"--encoder-bootstrap-port",
type=int,
default=ServerArgs.encoder_bootstrap_port,
help="Port for the EncoderBootstrapServer that runs in the "
"language-only tokenizer manager process. Encoders register here, "
"and language-only receivers fetch the current URL list from here.",
)
parser.add_argument(
"--encoder-register-urls",
nargs="+",
type=str,
default=[],
help="One or more EncoderBootstrapServer URLs to register this encoder "
"with on startup, for dynamic encoder discovery. "
"Example: --encoder-register-urls http://prefill0:8997 http://prefill1:8997. "
"Used with --encoder-only servers.",
)
parser.add_argument(
"--enable-adaptive-dispatch-to-encoder",
default=ServerArgs.enable_adaptive_dispatch_to_encoder,