Skip to content
Merged
Show file tree
Hide file tree
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 Apr 11, 2026
ebd94ea
test(mlx): remove unit tests to match vllm/sglang convention
key4ng Apr 11, 2026
6c60b06
test(mlx): add E2E test for MLX gRPC backend
key4ng Apr 11, 2026
a827734
ci(mlx): add GitHub workflow for MLX E2E tests on Apple Silicon
key4ng Apr 11, 2026
6ee5cce
ci(mlx): use debug maturin build to speed up workflow
key4ng Apr 11, 2026
17389ce
ci(mlx): use maturin build + pip install (no virtualenv) and fix ruff…
key4ng Apr 11, 2026
f202dc2
ci(mlx): use ci profile for maturin build
key4ng Apr 11, 2026
13c1520
fix(mlx): address PR #1099 review feedback
key4ng Apr 13, 2026
541a85d
fix(mlx): aggregate non-streaming logprobs + guard gRPC bind failure
key4ng Apr 13, 2026
024d542
test(mlx): hoist json import to module level
key4ng Apr 13, 2026
0e2ff44
refactor(mlx): hoist mlx.core and generation_stream imports to module…
key4ng Apr 13, 2026
b6fd962
fix(mlx): add threading.Lock around BatchGenerator state
key4ng Apr 13, 2026
eb35371
style(mlx): apply ruff format from pre-commit
key4ng Apr 13, 2026
d9c7d26
fix(mlx): address 3 new review comments on #1099
key4ng Apr 13, 2026
9910275
fix(mlx): fall back to 256 when no context-length key found
key4ng Apr 13, 2026
c2628d5
refactor(mlx): simplify servicer after review round
key4ng Apr 18, 2026
95866f6
fix(mlx): address two new codex review comments
key4ng Apr 18, 2026
0159b2b
fix(mlx): include chat template sidecars in tokenizer export
key4ng Apr 22, 2026
000f74c
chore(mlx): split E2E tests and CI workflow into separate branch
key4ng Apr 22, 2026
a05ac9c
chore(mlx): remove unused __main__.py
key4ng Apr 22, 2026
2d3638e
refactor(mlx): hoist remaining inline imports to module top
key4ng Apr 22, 2026
3b47bf5
fix(mlx): address 4 new codex review comments
key4ng Apr 23, 2026
5c24ce2
fix(mlx): report real backend status in HealthCheck
key4ng Apr 24, 2026
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
4 changes: 2 additions & 2 deletions crates/grpc_client/python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@ build-backend = "setuptools.build_meta"

[project]
name = "smg-grpc-proto"
version = "0.4.6"
description = "SMG gRPC proto definitions for SGLang, vLLM, and TRT-LLM"
version = "0.4.7"
description = "SMG gRPC proto definitions for SGLang, vLLM, TRT-LLM, and MLX"
requires-python = ">=3.10"
dependencies = [
"grpcio>=1.78.0",
Expand Down
6 changes: 5 additions & 1 deletion crates/grpc_client/python/smg_grpc_proto/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
"""SMG gRPC Proto - Protocol definitions for SGLang, vLLM, and TRT-LLM."""
"""SMG gRPC Proto - Protocol definitions for SGLang, vLLM, TRT-LLM, and MLX."""

from importlib.metadata import version

Expand All @@ -8,6 +8,8 @@
# These imports will work after the package is built (stubs generated at build time)
try:
from smg_grpc_proto.generated import (
mlx_engine_pb2,
mlx_engine_pb2_grpc,
sglang_encoder_pb2,
sglang_encoder_pb2_grpc,
sglang_scheduler_pb2,
Expand All @@ -27,6 +29,8 @@
"vllm_engine_pb2_grpc",
"trtllm_service_pb2",
"trtllm_service_pb2_grpc",
"mlx_engine_pb2",
"mlx_engine_pb2_grpc",
]
except ImportError:
# During development/build, generated modules may not exist yet
Expand Down
6 changes: 5 additions & 1 deletion grpc_servicer/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "smg-grpc-servicer"
version = "0.5.2"
description = "SMG gRPC servicer implementations for LLM inference engines (vLLM, SGLang)"
description = "SMG gRPC servicer implementations for LLM inference engines (vLLM, SGLang, MLX)"
requires-python = ">=3.10"
dependencies = [
"smg-grpc-proto>=0.4.6",
Expand All @@ -32,6 +32,10 @@ classifiers = [
[project.optional-dependencies]
vllm = ["vllm>=0.19.0"]
sglang = ["sglang>=0.5.10"]
# smg-grpc-proto>=0.4.7 is the first release that ships mlx_engine_pb2;
# without this floor, installing [mlx] against an older proto build would
# crash at import time when smg_grpc_servicer.mlx.server runs.
mlx = ["smg-grpc-proto>=0.4.7", "mlx>=0.22.0", "mlx-lm>=0.22.0"]

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Raise mlx-lm minimum version to BatchGenerator-capable release

mlx extra currently allows mlx-lm>=0.22.0, but this servicer imports BatchGenerator, SequenceStateMachine, and generation_stream from mlx_lm.generate (mlx/server.py and mlx/servicer.py). In mlx-lm 0.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 the mlx-lm floor to the first version that exports this batching/state-machine API.

Useful? React with 👍 / 👎.


[project.urls]
Homepage = "https://github.com/lightseekorg/smg"
Expand Down
6 changes: 6 additions & 0 deletions grpc_servicer/smg_grpc_servicer/mlx/__init__.py
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"]
69 changes: 69 additions & 0 deletions grpc_servicer/smg_grpc_servicer/mlx/health_servicer.py
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)
204 changes: 204 additions & 0 deletions grpc_servicer/smg_grpc_servicer/mlx/server.py
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*",
Comment thread
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()
Loading
Loading