From 04197fe27bbac27a65f0f59c78bb05b3fbf6449d Mon Sep 17 00:00:00 2001 From: Ryuichi Leo Takashige Date: Wed, 18 Mar 2026 21:21:48 +0000 Subject: [PATCH] Some QOL --- src/exo/disaggregated/batch_connector.py | 8 ++ src/exo/disaggregated/prefill_client.py | 35 ++++++-- src/exo/disaggregated/prefill_server.py | 82 ++++++++--------- src/exo/disaggregated/streaming_connector.py | 16 +++- src/exo/master/api.py | 2 +- src/exo/master/main.py | 4 +- src/exo/utils/info_gatherer/system_info.py | 87 +++++++++++++++---- .../engines/mlx/generator/batch_generate.py | 10 ++- src/exo/worker/engines/mlx/utils_mlx.py | 4 +- src/exo/worker/engines/vllm/vllm_generator.py | 18 +++- src/exo/worker/runner/llm_inference/runner.py | 10 +++ 11 files changed, 195 insertions(+), 81 deletions(-) diff --git a/src/exo/disaggregated/batch_connector.py b/src/exo/disaggregated/batch_connector.py index 2bb03c1d..b6c31c83 100644 --- a/src/exo/disaggregated/batch_connector.py +++ b/src/exo/disaggregated/batch_connector.py @@ -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 diff --git a/src/exo/disaggregated/prefill_client.py b/src/exo/disaggregated/prefill_client.py index 44562e93..068ba0be 100644 --- a/src/exo/disaggregated/prefill_client.py +++ b/src/exo/disaggregated/prefill_client.py @@ -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] diff --git a/src/exo/disaggregated/prefill_server.py b/src/exo/disaggregated/prefill_server.py index 24ed896f..746e3a6f 100644 --- a/src/exo/disaggregated/prefill_server.py +++ b/src/exo/disaggregated/prefill_server.py @@ -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] diff --git a/src/exo/disaggregated/streaming_connector.py b/src/exo/disaggregated/streaming_connector.py index 8102d7c0..24bcf7c2 100644 --- a/src/exo/disaggregated/streaming_connector.py +++ b/src/exo/disaggregated/streaming_connector.py @@ -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]: diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 50ea3248..3810f743 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -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) diff --git a/src/exo/master/main.py b/src/exo/master/main.py index 05d66d83..6726ccb2 100644 --- a/src/exo/master/main.py +++ b/src/exo/master/main.py @@ -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( diff --git a/src/exo/utils/info_gatherer/system_info.py b/src/exo/utils/info_gatherer/system_info.py index 2e9f619c..d360a679 100644 --- a/src/exo/utils/info_gatherer/system_info.py +++ b/src/exo/utils/info_gatherer/system_info.py @@ -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: diff --git a/src/exo/worker/engines/mlx/generator/batch_generate.py b/src/exo/worker/engines/mlx/generator/batch_generate.py index 11001a65..4fb94c91 100644 --- a/src/exo/worker/engines/mlx/generator/batch_generate.py +++ b/src/exo/worker/engines/mlx/generator/batch_generate.py @@ -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], diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index 516c073c..f84395a2 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -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 diff --git a/src/exo/worker/engines/vllm/vllm_generator.py b/src/exo/worker/engines/vllm/vllm_generator.py index 812b4401..8705bc55 100644 --- a/src/exo/worker/engines/vllm/vllm_generator.py +++ b/src/exo/worker/engines/vllm/vllm_generator.py @@ -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) diff --git a/src/exo/worker/runner/llm_inference/runner.py b/src/exo/worker/runner/llm_inference/runner.py index 375b0071..16bbb1d7 100644 --- a/src/exo/worker/runner/llm_inference/runner.py +++ b/src/exo/worker/runner/llm_inference/runner.py @@ -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: