This commit is contained in:
Ryuichi Leo Takashige
2026-03-18 23:12:38 +00:00
parent d1490444a1
commit 04197fe27b
11 changed files with 195 additions and 81 deletions
+8
View File
@@ -13,6 +13,12 @@ from vllm.distributed.kv_transfer.kv_connector.v1.base import ( # pyright: igno
_LAYER_RE = re.compile(r"layers\.(\d+)\.")
_active_instance: BatchConnector | None = None
def get_active_batch_connector() -> BatchConnector | None:
return _active_instance
@dataclass
class BatchConnectorMetadata(KVConnectorMetadata): # pyright: ignore[reportUntypedBaseClass]
@@ -25,6 +31,8 @@ class BatchConnector(KVConnectorBase_V1): # pyright: ignore[reportUntypedBaseCl
def __init__(self, vllm_config: Any, role: KVConnectorRole, kv_cache_config: Any = None) -> None: # type: ignore
super().__init__(vllm_config, role, kv_cache_config) # pyright: ignore[reportUnknownMemberType]
self.captured_layers = {}
global _active_instance
_active_instance = self
def start_load_kv(self, forward_context: Any, **kwargs: Any) -> None: # pyright: ignore[reportAny]
pass
+28 -7
View File
@@ -4,6 +4,7 @@ import json
import socket
import time
from collections import defaultdict
from collections.abc import Callable
from typing import TYPE_CHECKING, BinaryIO, cast
import mlx.core as mx
@@ -56,14 +57,19 @@ def _inject_rotating_kv_cache(cache: RotatingKVCache, keys: torch.Tensor, values
else:
keep = cache.keep
window = cache.max_size
sink_keys = k_mx[:, :, :keep, :]
sink_values = v_mx[:, :, :keep, :]
recent_keys = k_mx[:, :, -(window - keep):, :]
recent_values = v_mx[:, :, -(window - keep):, :]
cache.keys = mx.concatenate([sink_keys, recent_keys], axis=2)
cache.values = mx.concatenate([sink_values, recent_values], axis=2)
if keep == 0:
cache.keys = k_mx[:, :, -window:, :]
cache.values = v_mx[:, :, -window:, :]
cache._idx = window
else:
sink_keys = k_mx[:, :, :keep, :]
sink_values = v_mx[:, :, :keep, :]
recent_keys = k_mx[:, :, -(window - keep):, :]
recent_values = v_mx[:, :, -(window - keep):, :]
cache.keys = mx.concatenate([sink_keys, recent_keys], axis=2)
cache.values = mx.concatenate([sink_values, recent_values], axis=2)
cache._idx = keep
cache.offset = num_tokens
cache._idx = keep
def _inject_arrays_cache(cache: ArraysCache, arrays: list[torch.Tensor]) -> None:
@@ -75,6 +81,7 @@ def remote_prefill(
token_ids: list[int],
model_id: str,
mlx_model: Model,
on_prefill_progress: Callable[[int, int], None] | None = None,
) -> tuple[list[KVCache | RotatingKVCache | ArraysCache], int]:
if ":" in endpoint:
host, port_str = endpoint.rsplit(":", 1)
@@ -101,11 +108,16 @@ def remote_prefill(
raise RuntimeError(f"Prefill server error: {error_resp.get('error', 'unknown')}")
header = read_header(stream)
num_layers: int = header["num_layers"] # pyright: ignore[reportAssignmentType]
total_prompt_tokens = len(token_ids)
kv_buffers: dict[int, list[tuple[torch.Tensor, torch.Tensor]]] = defaultdict(list)
arrays_buffers: dict[int, list[torch.Tensor]] = {}
total_tokens = 0
layers_seen: set[int] = set()
tokens_received = 0
chunks_received = 0
t_first_chunk = None
while True:
msg = read_message(stream, header)
@@ -116,6 +128,15 @@ def remote_prefill(
if t_first_chunk is None:
t_first_chunk = time.perf_counter()
kv_buffers[msg.layer_idx].append((msg.keys, msg.values))
chunks_received += 1
layers_seen.add(msg.layer_idx)
tokens_received += msg.num_tokens
if on_prefill_progress and num_layers > 0 and chunks_received % num_layers == 0:
step = chunks_received // num_layers
on_prefill_progress(
min(tokens_received // num_layers, total_prompt_tokens),
total_prompt_tokens,
)
elif isinstance(msg, ArraysState):
arrays_buffers[msg.layer_idx] = msg.arrays
elif isinstance(msg, Done): # pyright: ignore[reportUnnecessaryIsInstance]
+35 -47
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import contextlib
import json
import queue
import socketserver
import threading
import time
@@ -97,21 +98,11 @@ def _run_prefill_overlapping(engine: LLMEngine, token_ids: list[int], wfile: Any
model_runner = get_model_runner()
assert model_runner is not None
connector: StreamingConnector | None = None
try:
engine_core = engine.engine_core.engine_core # type: ignore
scheduler = engine_core.scheduler # type: ignore
kv_manager = scheduler.kv_cache_manager # type: ignore
connector_obj = getattr(kv_manager, "connector", None) or getattr(scheduler, "connector", None) # pyright: ignore[reportUnknownArgumentType]
if isinstance(connector_obj, StreamingConnector):
connector = connector_obj
except Exception:
pass
from exo.disaggregated.streaming_connector import get_shared_queue, reset_shared_queue
if connector is None:
logger.warning("Could not find StreamingConnector, falling back to non-overlapping")
_run_prefill_batch(engine, token_ids, wfile)
return
reset_shared_queue()
layer_queue = get_shared_queue()
logger.info(f"Overlapping prefill: server reading from queue_id={id(layer_queue)}")
num_layers, dtype_str, layers_info = _get_layer_info(engine)
write_header(wfile, {"num_layers": num_layers, "dtype": dtype_str, "layers": layers_info}) # pyright: ignore[reportAny]
@@ -124,33 +115,33 @@ def _run_prefill_overlapping(engine: LLMEngine, token_ids: list[int], wfile: Any
params = SamplingParams(max_tokens=1, detokenize=False) # pyright: ignore[reportCallIssue]
engine.add_request(request_id, {"prompt_token_ids": token_ids}, params) # pyright: ignore[reportArgumentType]
prefill_done = threading.Event()
chunks_sent = [0]
def engine_loop() -> None:
while engine.has_unfinished_requests():
outputs = engine.step()
for output in outputs:
if output.request_id == request_id and output.outputs[0].token_ids:
engine.abort_request([request_id]) # type: ignore
connector.finish()
prefill_done.set()
return
connector.finish()
prefill_done.set()
def writer_loop() -> None:
while True:
item = layer_queue.get()
if item is None:
break
layer_idx, keys, values = item
write_kv_chunk(wfile, layer_idx, keys, values) # pyright: ignore[reportAny]
chunks_sent[0] += 1
engine_thread = threading.Thread(target=engine_loop, daemon=True)
engine_thread.start()
writer_thread = threading.Thread(target=writer_loop, daemon=True)
writer_thread.start()
layer_queue = connector.layer_queue
while True:
item = layer_queue.get()
if item is None:
break
layer_idx, keys, values = item
write_kv_chunk(wfile, layer_idx, keys, values) # pyright: ignore[reportAny]
while engine.has_unfinished_requests():
outputs = engine.step()
for output in outputs:
if output.request_id == request_id and output.outputs[0].token_ids:
engine.abort_request([request_id]) # type: ignore
break
else:
continue
break
prefill_done.wait()
engine_thread.join(timeout=5.0)
layer_queue.put(None)
writer_thread.join()
logger.info(f"Overlapping prefill: sent {chunks_sent[0]} KV chunks")
_stream_gdn_states(engine, wfile, num_layers, layers_info)
write_done(wfile, len(token_ids)) # pyright: ignore[reportAny]
@@ -165,16 +156,9 @@ def _run_prefill_batch(engine: LLMEngine, token_ids: list[int], wfile: Any) -> N
model_runner = get_model_runner()
assert model_runner is not None
connector: BatchConnector | None = None
try:
engine_core = engine.engine_core.engine_core # type: ignore
scheduler = engine_core.scheduler # type: ignore
kv_manager = scheduler.kv_cache_manager # type: ignore
connector_obj = getattr(kv_manager, "connector", None) or getattr(scheduler, "connector", None) # pyright: ignore[reportUnknownArgumentType]
if isinstance(connector_obj, BatchConnector):
connector = connector_obj
except Exception:
pass
from exo.disaggregated.batch_connector import get_active_batch_connector
connector = get_active_batch_connector()
from vllm.sampling_params import (
SamplingParams,
@@ -197,9 +181,13 @@ def _run_prefill_batch(engine: LLMEngine, token_ids: list[int], wfile: Any) -> N
write_header(wfile, {"num_layers": num_layers, "dtype": dtype_str, "layers": layers_info}) # pyright: ignore[reportAny]
if connector is not None:
logger.info(f"Batch prefill: streaming {len(connector.captured_layers)} captured layers")
for layer_idx in sorted(connector.captured_layers.keys()):
layer_data = connector.captured_layers[layer_idx]
write_kv_chunk(wfile, layer_idx, layer_data["keys"], layer_data["values"]) # pyright: ignore[reportAny]
connector.captured_layers.clear()
else:
logger.info("Batch prefill: no connector, sending 0 KV chunks")
_stream_gdn_states(engine, wfile, num_layers, layers_info)
write_done(wfile, len(token_ids)) # pyright: ignore[reportAny]
+15 -1
View File
@@ -14,6 +14,20 @@ from vllm.distributed.kv_transfer.kv_connector.v1.base import ( # pyright: igno
_LAYER_RE = re.compile(r"layers\.(\d+)\.")
_shared_queue: queue.Queue[tuple[int, torch.Tensor, torch.Tensor] | None] = queue.Queue()
def get_shared_queue() -> queue.Queue[tuple[int, torch.Tensor, torch.Tensor] | None]:
return _shared_queue
def reset_shared_queue() -> None:
while not _shared_queue.empty():
try:
_shared_queue.get_nowait()
except queue.Empty:
break
@dataclass
class StreamingConnectorMetadata(KVConnectorMetadata): # pyright: ignore[reportUntypedBaseClass]
@@ -25,7 +39,7 @@ class StreamingConnector(KVConnectorBase_V1): # pyright: ignore[reportUntypedBa
def __init__(self, vllm_config: Any, role: KVConnectorRole, kv_cache_config: Any = None) -> None: # type: ignore
super().__init__(vllm_config, role, kv_cache_config) # pyright: ignore[reportUnknownMemberType]
self._queue = queue.Queue()
self._queue = _shared_queue
@property
def layer_queue(self) -> queue.Queue[tuple[int, torch.Tensor, torch.Tensor] | None]:
+1 -1
View File
@@ -760,7 +760,7 @@ class API:
request_base = derive_base_model(str(model_id))
for instance in self.state.instances.values():
first_shard = next(iter(instance.shard_assignments.runner_to_shard.values()), None)
if first_shard is not None and first_shard.model_card.base_model == request_base:
if first_shard is not None and first_shard.model_card.base_model.lower() == request_base.lower():
return instance.shard_assignments.model_id
await self._trigger_notify_user_to_download_model(model_id)
+2 -2
View File
@@ -111,7 +111,7 @@ class Master:
if first_shard is None:
logger.info(f"Prefill routing: VllmInstance {instance.instance_id} has no shards")
continue
if first_shard.model_card.base_model != decode_model_base:
if first_shard.model_card.base_model.lower() != decode_model_base.lower():
logger.info(
f"Prefill routing: VllmInstance {instance.instance_id} base_model "
f"{first_shard.model_card.base_model!r} != decode {decode_model_base!r}"
@@ -203,7 +203,7 @@ class Master:
for instance in self.state.instances.values():
exact_match = instance.shard_assignments.model_id == command.task_params.model
first_shard = next(iter(instance.shard_assignments.runner_to_shard.values()), None)
base_match = first_shard is not None and first_shard.model_card.base_model == request_base
base_match = first_shard is not None and first_shard.model_card.base_model.lower() == request_base.lower()
if not (exact_match or base_match):
continue
task_count = sum(
+69 -18
View File
@@ -91,31 +91,82 @@ async def _get_interface_types_from_networksetup() -> dict[str, InterfaceType]:
return types
def _classify_unknown_darwin_interface(name: str) -> InterfaceType:
if name.lower().startswith("anpi"):
return "thunderbolt"
return "unknown"
async def _get_linux_network_interfaces() -> list[NetworkInterfaceInfo]:
import json as _json
try:
result = await run_process(["ip", "-j", "addr", "show"])
except (CalledProcessError, FileNotFoundError):
return []
data: list[dict[str, object]] = _json.loads(result.stdout) # pyright: ignore[reportAny]
interfaces: list[NetworkInterfaceInfo] = []
for iface in data:
name: str = iface.get("ifname", "") # pyright: ignore[reportAssignmentType, reportAny]
link_type: str = iface.get("link_type", "") # pyright: ignore[reportAssignmentType, reportAny]
iface_type: InterfaceType
if link_type == "loopback":
continue
elif link_type == "ether":
if name.startswith(("wl", "wlan")):
iface_type = "wifi"
elif name.startswith(("docker", "br-", "veth")):
iface_type = "unknown"
elif name.startswith(("thunderbolt", "tb", "enx")):
iface_type = "thunderbolt"
else:
iface_type = "ethernet"
elif link_type in ("none", "tun"):
iface_type = "unknown"
else:
iface_type = "unknown"
for addr_info in iface.get("addr_info", []): # pyright: ignore[reportAny]
family: str = addr_info.get("family", "") # pyright: ignore[reportAny]
ip: str = addr_info.get("local", "") # pyright: ignore[reportAny]
if family in ("inet", "inet6") and ip:
interfaces.append(NetworkInterfaceInfo(name=name, ip_address=ip, interface_type=iface_type))
return interfaces
async def get_network_interfaces() -> list[NetworkInterfaceInfo]:
"""
Retrieves detailed network interface information on macOS.
Parses output from 'networksetup -listallhardwareports' and 'ifconfig'
Retrieves detailed network interface information on macOS or Linux.
On MacOS: parses output from 'networksetup -listallhardwareports' and 'ifconfig'
to determine interface names, IP addresses, and types (ethernet, wifi, vpn, other).
Falls back to using ip -j addr show on other platforms.
Returns a list of NetworkInterfaceInfo objects.
"""
interfaces_info: list[NetworkInterfaceInfo] = []
interface_types = await _get_interface_types_from_networksetup()
for iface, services in psutil.net_if_addrs().items():
for service in services:
match service.family:
case socket.AF_INET | socket.AF_INET6:
interfaces_info.append(
NetworkInterfaceInfo(
name=iface,
ip_address=service.address,
interface_type=interface_types.get(iface, "unknown"),
if sys.platform == "darwin":
interfaces_info: list[NetworkInterfaceInfo] = []
interface_types = await _get_interface_types_from_networksetup()
for iface, services in psutil.net_if_addrs().items():
for service in services:
match service.family:
case socket.AF_INET | socket.AF_INET6:
iface_type = interface_types.get(iface, "unknown")
if iface_type == "unknown":
iface_type = _classify_unknown_darwin_interface(iface)
interfaces_info.append(
NetworkInterfaceInfo(
name=iface,
ip_address=service.address,
interface_type=iface_type,
)
)
)
case _:
pass
case _:
pass
return interfaces_info
return interfaces_info
return await _get_linux_network_interfaces()
def _read_dmi_field(name: str) -> str | None:
@@ -6,7 +6,7 @@ import mlx.core as mx
from mlx_lm.generate import (
BatchGenerator as MlxBatchGenerator,
)
from mlx_lm.models.cache import RotatingKVCache
from mlx_lm.models.cache import KVCache, RotatingKVCache
from mlx_lm.sample_utils import make_logits_processors, make_sampler
from mlx_lm.tokenizer_utils import StreamingDetokenizer, TokenizerWrapper
@@ -167,6 +167,7 @@ class ExoBatchGenerator:
token_ids=[int(t) for t in all_prompt_tokens.tolist()], # type: ignore
model_id=str(task_params.model),
mlx_model=self.model,
on_prefill_progress=on_prefill_progress,
)
cache = injected_cache
_prefill_tps = total_tokens / max(time.perf_counter() - t0, 0.001)
@@ -224,6 +225,13 @@ class ExoBatchGenerator:
max_tokens = task_params.max_output_tokens or MAX_TOKENS
if used_remote_prefill:
for ci, c in enumerate(cache):
if isinstance(c, RotatingKVCache):
logger.info(f"Cache[{ci}] RotatingKV: keys={c.keys.shape if c.keys is not None else None} _idx={c._idx} offset={c.offset} max_size={c.max_size}")
elif isinstance(c, KVCache):
logger.info(f"Cache[{ci}] KV: keys={c.keys.shape if c.keys is not None else None} offset={c.offset}")
uids = self._mlx_gen.insert(
prompts=[last_tokens.tolist()],
max_tokens=[max_tokens],
+2 -2
View File
@@ -551,7 +551,7 @@ def apply_chat_template(
)
if partial_assistant_content:
prompt += partial_assistant_content
logger.info(prompt)
logger.debug(prompt)
return prompt
extra_kwargs: dict[str, Any] = {}
@@ -588,7 +588,7 @@ def apply_chat_template(
if partial_assistant_content:
prompt += partial_assistant_content
logger.info(prompt)
logger.debug(prompt)
return prompt
+16 -2
View File
@@ -206,7 +206,7 @@ def vllm_generate(
on_generation_token: Callable[[], None] | None = None,
) -> Generator[GenerationResponse, None, None]:
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(engine, task)
logger.info(prompt_text)
logger.debug(prompt_text)
request_id = f"vllm-seq-{time.monotonic_ns()}"
sampling_params = make_vllm_sampling_params(engine, task, model_id)
engine.add_request(request_id, {"prompt_token_ids": token_ids}, sampling_params)
@@ -334,7 +334,7 @@ class VllmBatchEngine:
token_ids, prompt_text, prompt_token_count = format_vllm_prompt(
self.engine, task_params
)
logger.info(prompt_text)
logger.debug(prompt_text)
sampling_params = make_vllm_sampling_params(
self.engine, task_params, self.model_id
)
@@ -554,16 +554,29 @@ def load_vllm_engine(
trust_remote_code: bool,
n_layers: int = 1,
on_layer_loaded: Callable[[int, int], None] | None = None,
kv_connector_cls: type[object] | None = None,
) -> tuple[LLMEngine, ToolParser | None, KVPrefixCache]:
patch_vllm()
_patch_weight_loading_progress()
if kv_connector_cls is not None:
from exo.disaggregated.prefill_server import _patch_vllm_for_connector
_patch_vllm_for_connector(kv_connector_cls)
os.environ.setdefault("FASTSAFETENSORS_NOGDS", "1")
prefix_cache = KVPrefixCache(group=None)
set_prefix_cache(prefix_cache)
set_n_layers(n_layers)
kv_transfer_config: dict[str, str] | None = None
if kv_connector_cls is not None:
kv_transfer_config = {
"kv_connector": f"{kv_connector_cls.__module__}:{kv_connector_cls.__name__}",
"kv_role": "kv_both",
}
engine_args = EngineArgs(
model=model_path,
served_model_name=str(model_id),
@@ -574,6 +587,7 @@ def load_vllm_engine(
attention_backend="TRITON_ATTN",
enforce_eager=True,
disable_log_stats=True,
kv_transfer_config=kv_transfer_config, # type: ignore
)
set_weight_loading_callback(on_layer_loaded)
@@ -532,12 +532,22 @@ class VllmBuilder(Builder):
) -> None:
from exo.worker.engines.vllm.vllm_generator import load_vllm_engine
kv_connector_cls: type[object] | None = None
overlapping = not os.environ.get("EXO_NO_OVERLAPPING_PREFILL_SENDS")
if overlapping:
from exo.disaggregated.streaming_connector import StreamingConnector
kv_connector_cls = StreamingConnector
else:
from exo.disaggregated.batch_connector import BatchConnector
kv_connector_cls = BatchConnector
self._engine, self._tool_parser, self._prefix_cache = load_vllm_engine(
model_path=self.model_path,
model_id=self.model_id,
trust_remote_code=self.trust_remote_code,
n_layers=bound_instance.bound_shard.model_card.n_layers,
on_layer_loaded=on_layer_loaded,
kv_connector_cls=kv_connector_cls,
)
def build(self) -> InferenceGenerator: