Some QOL
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user