Swift-1.5-5bit-MLX / compatibility /swift15-server-cache.patch
ukisai's picture
Bound retained server cache while keeping prompt reuse enabled (#1)
8aff72b
Raw History Blame Contribute Delete
9.36 kB
diff --git a/mlx_lm/server.py b/mlx_lm/server.py
index 9462d6e..36f7787 100644
--- a/mlx_lm/server.py
+++ b/mlx_lm/server.py
@@ -29,6 +29,7 @@ from typing import (
import mlx.core as mx
from huggingface_hub import scan_cache_dir
+from mlx.utils import tree_flatten
from ._version import __version__
from .generate import (
@@ -428,6 +429,26 @@ def _format_top_logprobs(logprobs, top_n, tokenizer) -> Tuple[Dict[str, Any]]:
)
+def _prompt_cache_byte_limit(cli_args, model):
+ limit = getattr(cli_args, "prompt_cache_bytes", None)
+ if limit is not None:
+ if limit < 0:
+ raise ValueError("--prompt-cache-bytes must be non-negative")
+ return limit
+
+ # Reserve workspace and half the remaining capacity for the active request.
+ try:
+ recommended = mx.device_info().get("max_recommended_working_set_size")
+ except (RuntimeError, ValueError):
+ recommended = None
+ if not recommended:
+ return 1 << 63
+ parameters = {id(p): p for _, p in tree_flatten(model.parameters())}
+ model_bytes = sum(p.nbytes for p in parameters.values())
+ reserve = max(4 * 1024**3, recommended // 8)
+ return max(0, (recommended - model_bytes - reserve) // 2)
+
+
class ResponseGenerator:
def __init__(self, model_provider: ModelProvider, prompt_cache: LRUPromptCache):
self.model_provider = model_provider
@@ -443,6 +464,13 @@ class ResponseGenerator:
self._generation_thread = Thread(target=self._run_generate)
self._generation_thread.start()
+ def _configure_prompt_cache(self, model):
+ limit = _prompt_cache_byte_limit(self.cli_args, model)
+ if self.prompt_cache.max_bytes != limit:
+ self.prompt_cache.max_bytes = limit
+ self.prompt_cache.trim_to(n_bytes=limit)
+ logging.info("Retained prompt-cache limit: %.2f GB", limit / 1e9)
+
def _run_generate(self):
try:
self._generate()
@@ -772,6 +800,7 @@ class ResponseGenerator:
model, tokenizer = self.model_provider.load(
args.model.model, args.model.adapter, args.model.draft
)
+ self._configure_prompt_cache(model)
except Exception as e:
rqueue.put(e)
continue
@@ -1762,7 +1791,13 @@ def run(
handler_class=APIHandler,
):
group = mx.distributed.init()
- prompt_cache = LRUPromptCache(model_provider.cli_args.prompt_cache_size)
+ cache_bytes = model_provider.cli_args.prompt_cache_bytes
+ if cache_bytes is not None and cache_bytes < 0:
+ raise ValueError("--prompt-cache-bytes must be non-negative")
+ prompt_cache = LRUPromptCache(
+ model_provider.cli_args.prompt_cache_size,
+ max_bytes=cache_bytes if cache_bytes is not None else 1 << 63,
+ )
response_generator = ResponseGenerator(model_provider, prompt_cache)
if group.rank() == 0:
_run_http_server(host, port, response_generator)
@@ -1875,31 +1910,32 @@ def main():
parser.add_argument(
"--decode-concurrency",
type=int,
- default=32,
+ default=1,
help="When a request is batchable then decode that many requests in parallel",
)
parser.add_argument(
"--prompt-concurrency",
type=int,
- default=8,
+ default=1,
help="When a request is batchable then process that many prompts in parallel",
)
parser.add_argument(
"--prefill-step-size",
type=int,
- default=2048,
- help="Step size for prefill processing (default: 2048)",
+ default=512,
+ help="Step size for prefill processing (default: 512)",
)
parser.add_argument(
"--prompt-cache-size",
type=int,
- default=10,
- help="Maximum number of distinct KV caches to hold in the prompt cache",
+ default=2,
+ help="Maximum retained prompt-cache entries (default: 2)",
)
parser.add_argument(
"--prompt-cache-bytes",
type=_parse_size,
- help="Maximum size in bytes of the KV caches",
+ help="Maximum retained prompt-cache bytes. Default: automatic on Metal. "
+ "This does not cap active-request or total process memory.",
)
parser.add_argument(
"--kv-bits",
diff --git a/tests/test_server_cache_budget.py b/tests/test_server_cache_budget.py
new file mode 100644
--- /dev/null
+++ b/tests/test_server_cache_budget.py
@@ -0,0 +1,132 @@
+# Copyright © 2026 Swift contributors.
+
+import sys
+from types import SimpleNamespace
+from unittest.mock import Mock
+
+import pytest
+
+from mlx_lm import server
+from mlx_lm.models.cache import LRUPromptCache
+
+
+class CacheState:
+ def __init__(self, nbytes):
+ self.nbytes = nbytes
+
+ def is_trimmable(self):
+ return False
+
+
+@pytest.mark.parametrize("limit", [0, 100])
+def test_server_factory_enforces_configured_bytes(monkeypatch, limit):
+ caches = []
+ provider = SimpleNamespace(
+ cli_args=SimpleNamespace(prompt_cache_size=10, prompt_cache_bytes=limit)
+ )
+ monkeypatch.setattr(
+ server, "ResponseGenerator", lambda provider, cache: caches.append(cache)
+ )
+ monkeypatch.setattr(server, "_run_http_server", lambda *args: None)
+ server.run("127.0.0.1", 0, provider)
+ cache = caches[0]
+ for i in range(5):
+ cache.insert_cache("model", [i, 1], [CacheState(80)])
+ assert cache.nbytes <= limit
+ if limit:
+ reused, remaining = cache.fetch_nearest_cache("model", [4, 1, 9])
+ assert reused is not None
+ assert remaining == [9]
+
+
+def test_negative_limit_rejected_before_worker_start(monkeypatch):
+ worker = Mock()
+ monkeypatch.setattr(server, "ResponseGenerator", worker)
+ provider = SimpleNamespace(
+ cli_args=SimpleNamespace(prompt_cache_size=2, prompt_cache_bytes=-1)
+ )
+ with pytest.raises(ValueError, match="non-negative"):
+ server.run("127.0.0.1", 0, provider)
+ worker.assert_not_called()
+
+
+def test_explicit_limit_does_not_probe_model(monkeypatch):
+ probe = Mock(side_effect=AssertionError("device probe was not needed"))
+ monkeypatch.setattr(server.mx, "device_info", probe)
+ model = Mock()
+ for limit in (0, 1024, 8 * 1024**3):
+ assert (
+ server._prompt_cache_byte_limit(
+ SimpleNamespace(prompt_cache_bytes=limit), model
+ )
+ == limit
+ )
+ model.parameters.assert_not_called()
+
+
+def test_auto_limit_reserves_room_for_active_request(monkeypatch):
+ gib = 1024**3
+ monkeypatch.setattr(
+ server.mx,
+ "device_info",
+ lambda: {"max_recommended_working_set_size": 36 * gib},
+ )
+ model = SimpleNamespace(parameters=lambda: {"weight": CacheState(15 * gib)})
+ args = SimpleNamespace(prompt_cache_bytes=None)
+ limit = server._prompt_cache_byte_limit(args, model)
+ assert 6 * gib <= limit < (36 - 15) * gib // 2
+ larger = SimpleNamespace(parameters=lambda: {"weight": CacheState(30 * gib)})
+ assert server._prompt_cache_byte_limit(args, larger) < limit
+ full = SimpleNamespace(parameters=lambda: {"weight": CacheState(36 * gib)})
+ assert server._prompt_cache_byte_limit(args, full) == 0
+
+
+def test_model_swap_trims_retained_states(monkeypatch):
+ obj = server.ResponseGenerator.__new__(server.ResponseGenerator)
+ obj.model_provider = SimpleNamespace(
+ cli_args=SimpleNamespace(prompt_cache_bytes=100)
+ )
+ obj.prompt_cache = LRUPromptCache(max_size=10)
+ obj.prompt_cache.insert_cache("old", [1], [CacheState(200)])
+ obj._configure_prompt_cache(Mock())
+ assert obj.prompt_cache.nbytes == 0
+ assert obj.prompt_cache.max_bytes == 100
+
+
+@pytest.mark.parametrize("device_info", [{}, {"max_recommended_working_set_size": 0}])
+def test_auto_limit_without_metal_metadata(monkeypatch, device_info):
+ monkeypatch.setattr(server.mx, "device_info", lambda: device_info)
+ model = Mock()
+ assert (
+ server._prompt_cache_byte_limit(SimpleNamespace(prompt_cache_bytes=None), model)
+ == 1 << 63
+ )
+ model.parameters.assert_not_called()
+
+
+def test_auto_limit_counts_tied_parameters_once(monkeypatch):
+ gib = 1024**3
+ monkeypatch.setattr(
+ server.mx,
+ "device_info",
+ lambda: {"max_recommended_working_set_size": 32 * gib},
+ )
+ weight = CacheState(8 * gib)
+ model = SimpleNamespace(parameters=lambda: {"embed": weight, "head": weight})
+ assert (
+ server._prompt_cache_byte_limit(SimpleNamespace(prompt_cache_bytes=None), model)
+ == 10 * gib
+ )
+
+
+def test_defaults_keep_reuse_enabled_and_limit_concurrency(monkeypatch):
+ captured = []
+ monkeypatch.setattr(sys, "argv", ["mlx_lm.server"])
+ monkeypatch.setattr(server, "maybe_set_recommended_wired_limit", lambda: None)
+ monkeypatch.setattr(server, "run", lambda h, p, m: captured.append(m.cli_args))
+ server.main()
+ args = captured[0]
+ assert args.prompt_cache_size == 2
+ assert args.prompt_cache_bytes is None
+ assert args.prompt_concurrency == args.decode_concurrency == 1
+ assert args.prefill_step_size == 512