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