-
Notifications
You must be signed in to change notification settings - Fork 179
feat(mlx): add Python gRPC servicer for MLX backend #1099
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
23 commits
Select commit
Hold shift + click to select a range
bf9da2c
feat(mlx): add Python servicer for MLX gRPC backend
key4ng ebd94ea
test(mlx): remove unit tests to match vllm/sglang convention
key4ng 6c60b06
test(mlx): add E2E test for MLX gRPC backend
key4ng a827734
ci(mlx): add GitHub workflow for MLX E2E tests on Apple Silicon
key4ng 6ee5cce
ci(mlx): use debug maturin build to speed up workflow
key4ng 17389ce
ci(mlx): use maturin build + pip install (no virtualenv) and fix ruff…
key4ng f202dc2
ci(mlx): use ci profile for maturin build
key4ng 13c1520
fix(mlx): address PR #1099 review feedback
key4ng 541a85d
fix(mlx): aggregate non-streaming logprobs + guard gRPC bind failure
key4ng 024d542
test(mlx): hoist json import to module level
key4ng 0e2ff44
refactor(mlx): hoist mlx.core and generation_stream imports to module…
key4ng b6fd962
fix(mlx): add threading.Lock around BatchGenerator state
key4ng eb35371
style(mlx): apply ruff format from pre-commit
key4ng d9c7d26
fix(mlx): address 3 new review comments on #1099
key4ng 9910275
fix(mlx): fall back to 256 when no context-length key found
key4ng c2628d5
refactor(mlx): simplify servicer after review round
key4ng 95866f6
fix(mlx): address two new codex review comments
key4ng 0159b2b
fix(mlx): include chat template sidecars in tokenizer export
key4ng 000f74c
chore(mlx): split E2E tests and CI workflow into separate branch
key4ng a05ac9c
chore(mlx): remove unused __main__.py
key4ng 2d3638e
refactor(mlx): hoist remaining inline imports to module top
key4ng 3b47bf5
fix(mlx): address 4 new codex review comments
key4ng 5c24ce2
fix(mlx): report real backend status in HealthCheck
key4ng File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,6 @@ | ||
| """MLX gRPC servicers -- MlxEngine proto service and standard health check.""" | ||
|
|
||
| from smg_grpc_servicer.mlx.health_servicer import MlxHealthServicer | ||
| from smg_grpc_servicer.mlx.servicer import MlxEngineServicer | ||
|
|
||
| __all__ = ["MlxEngineServicer", "MlxHealthServicer"] |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,69 @@ | ||
| """ | ||
| Standard gRPC health check service for MLX. | ||
|
|
||
| Implements grpc.health.v1.Health protocol with simple liveness tracking. | ||
| """ | ||
|
|
||
| import logging | ||
| from collections.abc import AsyncIterator | ||
|
|
||
| import grpc | ||
| from grpc_health.v1 import health_pb2, health_pb2_grpc | ||
|
|
||
| logger = logging.getLogger(__name__) | ||
|
|
||
|
|
||
| class MlxHealthServicer(health_pb2_grpc.HealthServicer): | ||
| """Standard gRPC health check for MLX inference server.""" | ||
|
|
||
| OVERALL_SERVER = "" | ||
| MLX_SERVICE = "mlx.grpc.engine.MlxEngine" | ||
|
|
||
| def __init__(self): | ||
| self._serving = False | ||
| logger.info("MlxHealthServicer initialized") | ||
|
|
||
| def set_serving(self): | ||
| self._serving = True | ||
| logger.info("Health status set to SERVING") | ||
|
|
||
| def set_not_serving(self): | ||
| self._serving = False | ||
| logger.info("Health status set to NOT_SERVING") | ||
|
|
||
| async def Check( | ||
| self, | ||
| request: health_pb2.HealthCheckRequest, | ||
| context: grpc.aio.ServicerContext, | ||
| ) -> health_pb2.HealthCheckResponse: | ||
| service_name = request.service | ||
|
|
||
| if service_name in (self.OVERALL_SERVER, self.MLX_SERVICE): | ||
| status = ( | ||
| health_pb2.HealthCheckResponse.SERVING | ||
| if self._serving | ||
| else health_pb2.HealthCheckResponse.NOT_SERVING | ||
| ) | ||
| return health_pb2.HealthCheckResponse(status=status) | ||
|
|
||
| context.set_code(grpc.StatusCode.NOT_FOUND) | ||
| context.set_details(f"Unknown service: {service_name}") | ||
| return health_pb2.HealthCheckResponse(status=health_pb2.HealthCheckResponse.SERVICE_UNKNOWN) | ||
|
|
||
| async def Watch( | ||
| self, | ||
| request: health_pb2.HealthCheckRequest, | ||
| context: grpc.aio.ServicerContext, | ||
| ) -> AsyncIterator[health_pb2.HealthCheckResponse]: | ||
| service_name = request.service | ||
|
|
||
| if service_name in (self.OVERALL_SERVER, self.MLX_SERVICE): | ||
| status = ( | ||
| health_pb2.HealthCheckResponse.SERVING | ||
| if self._serving | ||
| else health_pb2.HealthCheckResponse.NOT_SERVING | ||
| ) | ||
| else: | ||
| status = health_pb2.HealthCheckResponse.SERVICE_UNKNOWN | ||
|
|
||
| yield health_pb2.HealthCheckResponse(status=status) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,204 @@ | ||
| """ | ||
| MLX gRPC Server | ||
|
|
||
| Standalone gRPC server entrypoint for MLX inference. | ||
| CLI: python -m smg_grpc_servicer.mlx.server --model <path> --port 50051 | ||
| """ | ||
|
|
||
| import argparse | ||
| import asyncio | ||
| import json | ||
| import logging | ||
| import os | ||
| import signal | ||
| import time | ||
| from concurrent import futures | ||
|
|
||
| import grpc | ||
| from grpc_health.v1 import health_pb2_grpc | ||
| from grpc_reflection.v1alpha import reflection | ||
| from huggingface_hub import snapshot_download | ||
| from mlx_lm import load | ||
| from mlx_lm.generate import BatchGenerator | ||
| from smg_grpc_proto import mlx_engine_pb2, mlx_engine_pb2_grpc | ||
|
|
||
| from smg_grpc_servicer.mlx.health_servicer import MlxHealthServicer | ||
| from smg_grpc_servicer.mlx.servicer import MlxEngineServicer | ||
|
|
||
| logger = logging.getLogger(__name__) | ||
|
|
||
|
|
||
| def parse_args(): | ||
| parser = argparse.ArgumentParser(description="MLX gRPC inference server") | ||
| parser.add_argument("--model", required=True, help="Model path or HuggingFace repo ID") | ||
| parser.add_argument("--port", type=int, default=50051, help="gRPC listen port") | ||
| parser.add_argument("--host", default="0.0.0.0", help="gRPC listen address") | ||
| parser.add_argument( | ||
| "--prefill-batch-size", type=int, default=8, help="Max concurrent prefill requests" | ||
| ) | ||
| parser.add_argument( | ||
| "--completion-batch-size", type=int, default=32, help="Max concurrent generation requests" | ||
| ) | ||
| parser.add_argument("--adapter-path", default=None, help="LoRA adapter path") | ||
| return parser.parse_args() | ||
|
|
||
|
|
||
| def load_model(args): | ||
| """Load model and tokenizer via mlx-lm.""" | ||
| logger.info("Loading model: %s", args.model) | ||
| model, tokenizer = load(args.model, adapter_path=args.adapter_path) | ||
| logger.info("Model loaded successfully") | ||
|
|
||
| model_dir = args.model | ||
| if not os.path.isdir(model_dir): | ||
| model_dir = snapshot_download( | ||
| args.model, | ||
| allow_patterns=[ | ||
| "config.json", | ||
| "tokenizer*", | ||
| "special_tokens*", | ||
|
key4ng marked this conversation as resolved.
|
||
| "merges.txt", | ||
| "vocab.json", | ||
| "added_tokens.json", | ||
| # Chat template sidecars (Gemma 4, Llama 3.1+, newer models). | ||
| "chat_template.json", | ||
| "chat_template.jinja", | ||
| # tiktoken-style tokenizer artifacts — must stay in sync | ||
| # with MlxEngineServicer._TOKENIZER_FILES / _SUFFIXES. | ||
| "tiktoken.model", | ||
| "*.tiktoken", | ||
| ], | ||
| ) | ||
|
|
||
| config_path = os.path.join(model_dir, "config.json") | ||
| with open(config_path) as f: | ||
| model_config = json.load(f) | ||
|
|
||
| eos = model_config.get("eos_token_id") | ||
| if isinstance(eos, int): | ||
| eos_token_ids = [eos] | ||
| elif isinstance(eos, list): | ||
| eos_token_ids = eos | ||
| else: | ||
| eos_token_ids = list(tokenizer.eos_token_ids) if hasattr(tokenizer, "eos_token_ids") else [] | ||
|
|
||
| return model, tokenizer, model_dir, model_config, eos_token_ids | ||
|
|
||
|
|
||
| def _warmup(batch_generator): | ||
| """Run one end-to-end token through the batch generator so the first | ||
| real request doesn't pay JIT/kernel compilation cost.""" | ||
| logger.info("Running warmup generation...") | ||
| try: | ||
| uids = batch_generator.insert(prompts=[[1]], max_tokens=[1]) | ||
| for _ in range(10): | ||
| _, gen_responses = batch_generator.next() | ||
| if any(r.finish_reason is not None for r in gen_responses if r.uid == uids[0]): | ||
| break | ||
| batch_generator.remove(uids) | ||
| logger.info("Warmup complete") | ||
| except Exception: | ||
| logger.warning("Warmup failed (non-fatal)", exc_info=True) | ||
|
|
||
|
|
||
| async def serve_grpc(args): | ||
| """Start the MLX gRPC server.""" | ||
| start_time = time.time() | ||
|
|
||
| model, tokenizer, model_dir, model_config, eos_token_ids = load_model(args) | ||
|
|
||
| batch_generator = BatchGenerator( | ||
| model, | ||
| completion_batch_size=args.completion_batch_size, | ||
| prefill_batch_size=args.prefill_batch_size, | ||
| ) | ||
| logger.info( | ||
| "BatchGenerator created (prefill=%d, completion=%d)", | ||
| args.prefill_batch_size, | ||
| args.completion_batch_size, | ||
| ) | ||
|
|
||
| server = grpc.aio.server( | ||
| futures.ThreadPoolExecutor(max_workers=10), | ||
| options=[ | ||
| ("grpc.max_send_message_length", 1024 * 1024 * 256), | ||
| ("grpc.max_receive_message_length", 1024 * 1024 * 256), | ||
| ("grpc.http2.min_recv_ping_interval_without_data_ms", 10000), | ||
| ("grpc.keepalive_permit_without_calls", True), | ||
| ], | ||
| ) | ||
|
|
||
| health_servicer = MlxHealthServicer() | ||
| health_pb2_grpc.add_HealthServicer_to_server(health_servicer, server) | ||
|
|
||
| servicer = MlxEngineServicer( | ||
| batch_generator=batch_generator, | ||
| model_path=args.model, | ||
| model_dir=model_dir, | ||
| model_config=model_config, | ||
| eos_token_ids=eos_token_ids, | ||
| start_time=start_time, | ||
| ) | ||
| mlx_engine_pb2_grpc.add_MlxEngineServicer_to_server(servicer, server) | ||
|
|
||
| SERVICE_NAMES = ( | ||
| mlx_engine_pb2.DESCRIPTOR.services_by_name["MlxEngine"].full_name, | ||
| "grpc.health.v1.Health", | ||
| reflection.SERVICE_NAME, | ||
| ) | ||
| reflection.enable_server_reflection(SERVICE_NAMES, server) | ||
|
|
||
| listen_addr = f"{args.host}:{args.port}" | ||
| bound_port = server.add_insecure_port(listen_addr) | ||
| if bound_port == 0: | ||
| raise RuntimeError(f"Failed to bind gRPC server to {listen_addr}") | ||
|
|
||
| # Warmup BEFORE starting the generation loop (batch_generator.next() is | ||
| # not thread-safe — only one caller at a time). | ||
| _warmup(batch_generator) | ||
| servicer.start_generation_loop() | ||
|
|
||
| # Only accept RPCs after the generation loop is running. Otherwise a | ||
| # Generate RPC could slip into the window between server.start() and | ||
| # start_generation_loop() and block forever on queue.get() because no | ||
| # gen thread is dispatching tokens. HealthCheck always returns OK, so | ||
| # the router can't use it to detect this window. | ||
| await server.start() | ||
| health_servicer.set_serving() | ||
| logger.info("gRPC server listening on %s — model: %s", listen_addr, args.model) | ||
|
|
||
| loop = asyncio.get_running_loop() | ||
| stop_event = asyncio.Event() | ||
|
|
||
| def signal_handler(): | ||
| logger.info("Received shutdown signal") | ||
| stop_event.set() | ||
|
|
||
| for sig in (signal.SIGTERM, signal.SIGINT): | ||
| loop.add_signal_handler(sig, signal_handler) | ||
|
|
||
| try: | ||
| await stop_event.wait() | ||
| finally: | ||
| logger.info("Shutting down...") | ||
| health_servicer.set_not_serving() | ||
| # Stop accepting new RPCs first so in-flight requests can still | ||
| # drain against the running generation thread. Stopping the gen | ||
| # loop first would leave new/in-flight RPCs stranded. | ||
| await server.stop(5.0) | ||
| servicer.stop_generation_loop() | ||
| batch_generator.close() | ||
| logger.info("Server stopped") | ||
|
|
||
|
|
||
| def main(): | ||
| logging.basicConfig( | ||
| level=logging.INFO, | ||
| format="%(asctime)s %(levelname)s %(name)s: %(message)s", | ||
| ) | ||
| args = parse_args() | ||
| asyncio.run(serve_grpc(args)) | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| main() | ||
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
mlxextra currently allowsmlx-lm>=0.22.0, but this servicer importsBatchGenerator,SequenceStateMachine, andgeneration_streamfrommlx_lm.generate(mlx/server.pyandmlx/servicer.py). Inmlx-lm0.22.x those symbols are not available, so a resolver that pins 0.22.* (still valid under this constraint) will fail at import/startup time before serving any RPCs. Please bump themlx-lmfloor to the first version that exports this batching/state-machine API.Useful? React with 👍 / 👎.