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
63 changes: 30 additions & 33 deletions mlx_engine/model_kit/batched_model_kit.py
Original file line number Diff line number Diff line change
Expand Up @@ -334,10 +334,9 @@ def get_next_request(timeout=None):
timeout: None | float = None if (len(self._batch_results) > 0) else 0.1
request = get_next_request(timeout=timeout)

# We got a request
# Handle at most one request before advancing the current batch.
if request is not None:
if isinstance(request, CancelGenerationRequest):
# Handle cancel request
found_request_id = False
request_id = request.request_id
for uid, entry in self._batch_results.items():
Expand All @@ -349,45 +348,40 @@ def get_next_request(timeout=None):
break
if not found_request_id:
logger.warning(f"Could not cancel {request_id=} (id not found)")
continue

with mx.stream(batch_generator.stream):
cache, cached_prefix, rest = _prepare_prompt_cache_for_generation(
self._prompt_cache, current_model_key, request.prompt_tokens
)

# Keep cache allocation on the same MLX stream as generation.
(uid,) = batch_generator.insert(
[rest],
[request.max_tokens],
caches=[cache],
all_tokens=[cached_prefix],
samplers=[request.samplers],
logits_processors=[request.logits_processors],
)
else:
with mx.stream(batch_generator.stream):
cache, cached_prefix, rest = (
_prepare_prompt_cache_for_generation(
self._prompt_cache,
current_model_key,
request.prompt_tokens,
)
)

# Track this request
self._batch_results[uid] = {
"cache_key": request.prompt_tokens[:],
"rqueue": request.rqueue,
"detokenizer": self.tokenizer.detokenizer,
"top_logprobs": request.top_logprobs,
"request_id": request.request_id,
}
# Keep cache allocation on the same MLX stream as generation.
(uid,) = batch_generator.insert(
[rest],
[request.max_tokens],
caches=[cache],
all_tokens=[cached_prefix],
samplers=[request.samplers],
logits_processors=[request.logits_processors],
)

# Check for new requests
continue
self._batch_results[uid] = {
"cache_key": request.prompt_tokens[:],
"rqueue": request.rqueue,
"detokenizer": self.tokenizer.detokenizer,
"top_logprobs": request.top_logprobs,
"request_id": request.request_id,
}

# No request so serve from the current batch
if len(self._batch_results) == 0:
if self._shutdown.is_set() or len(self._batch_results) == 0:
continue

time_budget = 0.5
start = time.time()
while True:
if time.time() - start > time_budget:
break

prompt_responses, generation_responses = batch_generator.next()
if not prompt_responses and not generation_responses:
break
Expand Down Expand Up @@ -449,6 +443,9 @@ def get_next_request(timeout=None):
)
del self._batch_results[r.uid]

if not self._requests.empty() or time.time() - start > time_budget:
break

for entry in self._batch_results.values():
entry["rqueue"].put(RequestCancelled("Model shutdown requested"))

Expand Down
6 changes: 3 additions & 3 deletions mlx_engine/model_kit/batched_vision/model_kit.py
Original file line number Diff line number Diff line change
Expand Up @@ -512,12 +512,12 @@ def _insert_prepared_request(
)
detokenizer = self._new_detokenizer()

total_prompt_tokens = max(0, prompt_token_count - 1)
cached_tokens = min(cached_prefix_len, total_prompt_tokens)
prefill_tokens = max(0, prompt_token_count - 1)
cached_tokens = min(cached_prefix_len, prefill_tokens)
request.rqueue.put(
PromptProgressBeginEvent(
cached_tokens=cached_tokens,
total_prompt_tokens=total_prompt_tokens,
total_prompt_tokens=prompt_token_count,
prefill_tokens_processed=0,
)
)
Expand Down
5 changes: 5 additions & 0 deletions mlx_engine/server/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
"""Private HTTP runtime for mlx-engine."""

from .http import EngineRuntime, MlxEngineHttpServer

__all__ = ["EngineRuntime", "MlxEngineHttpServer"]
76 changes: 76 additions & 0 deletions mlx_engine/server/__main__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
import argparse
import logging
import os
import signal
import threading

from mlx_engine import load_model

from .http import EngineRuntime, MlxEngineHttpServer


logger = logging.getLogger(__name__)


_API_KEY_ENV_VAR = "MLX_ENGINE_API_KEY"


def _create_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Run the private mlx-engine server.",
epilog=f"Authentication is configured through {_API_KEY_ENV_VAR}.",
)
parser.add_argument("--model", required=True)
parser.add_argument("--host", required=True)
parser.add_argument("--port", required=True, type=int)
parser.add_argument("--context-length", required=True, type=int)
parser.add_argument("--parallel-sessions", required=True, type=int)
parser.add_argument("--seed", type=int)
return parser


def main() -> None:
parser = _create_parser()
args = parser.parse_args()
api_key = os.environ.get(_API_KEY_ENV_VAR)
if not api_key:
parser.error(f"{_API_KEY_ENV_VAR} environment variable is required")

logger.info("Loading MLX model from %s", args.model)
model_kit = load_model(
args.model,
max_kv_size=args.context_length,
max_seq_nums=args.parallel_sessions,
seed=args.seed,
trust_remote_code=False,
)
runtime = EngineRuntime(model_kit)
server = None

try:
server = MlxEngineHttpServer(
(args.host, args.port),
api_key=api_key,
runtime=runtime,
)

def request_shutdown(_signal_number: int, _frame: object) -> None:
logger.info("Stopping MLX server")
server.cancel_active_sessions()
threading.Thread(target=server.shutdown, daemon=True).start()

signal.signal(signal.SIGINT, request_shutdown)
signal.signal(signal.SIGTERM, request_shutdown)
logger.info("MLX server listening on %s:%d", args.host, args.port)
server.serve_forever()
finally:
try:
if server is not None:
server.cancel_active_sessions()
server.server_close()
finally:
runtime.unload()


if __name__ == "__main__":
main()
Loading
Loading