diff --git a/mlx_lm/evaluate.py b/mlx_lm/evaluate.py index aed76d83d..340eb69e6 100644 --- a/mlx_lm/evaluate.py +++ b/mlx_lm/evaluate.py @@ -26,7 +26,7 @@ from .generate import batch_generate from .models.cache import make_prompt_cache from .sample_utils import make_sampler -from .utils import common_prefix_len, load +from .utils import load DEFAULT_MAX_TOKENS = 8192 diff --git a/mlx_lm/server.py b/mlx_lm/server.py index 410c83a02..cca8101f9 100644 --- a/mlx_lm/server.py +++ b/mlx_lm/server.py @@ -1,6 +1,7 @@ # Copyright © 2023-2024 Apple Inc. import argparse +import copy import json import logging import platform @@ -8,6 +9,7 @@ import time import uuid import warnings +from collections import deque from dataclasses import dataclass, field from http.server import BaseHTTPRequestHandler, HTTPServer from pathlib import Path @@ -30,7 +32,7 @@ from .generate import stream_generate from .models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache from .sample_utils import make_logits_processors, make_sampler -from .utils import common_prefix_len, load +from .utils import load def get_system_fingerprint(): @@ -145,6 +147,143 @@ def process_message_content(messages): message["content"] = "" +class LRUPromptCache: + + @dataclass + class CacheEntry: + prompt_cache: List[Any] + count: int + + @dataclass + class SearchResult: + model: Any + exact: List[int] + shorter: List[int] + longer: List[int] + common_prefix: int + + def __init__(self, max_size: int = 10): + self.max_size = max_size + self._cache = {} + self._lru = deque() + + def _search(self, model, tokens): + """Search the cache for a prompt cache. Return exact or close match.""" + if model not in self._cache: + return self.SearchResult(model, None, None, None, 0) + + current = self._cache[model] + last_cache_index = -1 + index = 0 + + while index < len(tokens) and tokens[index] in current: + current = current[tokens[index]] + if "cache" in current: + last_cache_index = index + index += 1 + + # Exact match no need to search for longer or shorter caches + if last_cache_index == len(tokens) - 1: + return self.SearchResult(model, tokens, None, None, 0) + + # Find the shorter cache + shorter = None + if last_cache_index > 0: + shorter = tokens[: last_cache_index + 1] + + # Check for caches that are longer + longer = None + common_prefix = index + if index > 0 and last_cache_index <= 0: + best = None + stack = [(current, [])] + while stack: + current, extra = stack.pop() + if "cache" in current: + if best is None or len(extra) < len(best): + best = extra + else: + for tok in current: + stack.append((current[tok], extra + [tok])) + longer = tokens[:index] + best + return self.SearchResult(model, None, shorter, longer, common_prefix) + + def _get(self, model, tokens): + current = self._cache[model] + for tok in tokens: + current = current[tok] + return current["cache"] + + def _delete(self, model, tokens): + path = [self._cache[model]] + for tok in tokens: + path.append(path[-1][tok]) + del path[-1]["cache"] + for i in reversed(range(len(tokens))): + d_prev, d, t = path[i], path[i + 1], tokens[i] + if len(d) > 0: + break + del d_prev[t] + + def _extract(self, model, tokens): + cache_entry = self._get(model, tokens) + if cache_entry.count == 1: + self._delete(model, tokens) + self._lru.remove((model, tokens)) + return cache_entry + + cache_entry.count -= 1 + return self.CacheEntry( + copy.deepcopy(cache_entry.prompt_cache), + 1, + ) + + def fetch_nearest_cache(self, model, tokens): + result = self._search(model, tokens) + if result.exact is not None: + cache_entry = self._extract(result.model, result.exact) + return cache_entry.prompt_cache, [] + + if result.shorter is not None: + cache_entry = self._extract(result.model, result.shorter) + prefix_len = len(result.shorter) + return cache_entry.prompt_cache, tokens[prefix_len:] + + if result.longer is not None: + cache_entry = self._get(result.model, result.longer) + if can_trim_prompt_cache(cache_entry.prompt_cache): + cache_entry = self.CacheEntry( + copy.deepcopy(cache_entry.prompt_cache), + 1, + ) + prefix = min(len(tokens) - 1, result.common_prefix) + num_to_trim = len(result.longer) - prefix + trim_prompt_cache(cache_entry.prompt_cache, num_to_trim) + return cache_entry.prompt_cache, tokens[prefix:] + + return None, tokens + + def insert_cache(self, model, tokens, prompt_cache): + if model not in self._cache: + self._cache[model] = {} + current = self._cache[model] + for tok in tokens: + if tok not in current: + current[tok] = {} + current = current[tok] + + if "cache" in current: + current["cache"].count += 1 + self._lru.remove((model, tokens)) + else: + current["cache"] = self.CacheEntry(prompt_cache, 1) + + self._lru.append((model, tokens)) + if len(self._lru) > self.max_size: + model, tokens = self._lru.popleft() + self._delete(model, tokens) + + @dataclass class PromptCache: cache: List[Any] = field(default_factory=list) @@ -247,7 +386,7 @@ def __init__( """ self.created = int(time.time()) self.model_provider = model_provider - self.prompt_cache = prompt_cache or PromptCache() + self.prompt_cache = prompt_cache or LRUPromptCache() self.system_fingerprint = system_fingerprint or get_system_fingerprint() super().__init__(*args, **kwargs) @@ -538,88 +677,33 @@ def parse_function(tool_text): return response - def reset_prompt_cache(self, prompt): - """Resets the prompt cache and associated state. - - Args: - prompt (List[int]): The tokenized new prompt which will populate the - reset cache. - """ - logging.debug(f"*** Resetting cache. ***") - self.prompt_cache.model_key = self.model_provider.model_key - self.prompt_cache.cache = make_prompt_cache(self.model_provider.model) - if self.model_provider.draft_model is not None: - self.prompt_cache.cache += make_prompt_cache( - self.model_provider.draft_model - ) - self.prompt_cache.tokens = list(prompt) # Cache the new prompt fully - def get_prompt_cache(self, prompt): """ - Determines the portion of the prompt that needs processing by comparing - it to the cached prompt and attempting to reuse the common prefix. + Given the prompt find the closest KV cache that can be extended to the + passed in prompt. - This function updates the internal prompt cache state (tokens and model cache) - based on the comparison. If a common prefix exists, it attempts to trim - the model cache (if supported) to match the common prefix length, avoiding - recomputation. + If one couldn't be found then make a new one. Args: prompt (List[int]): The tokenized new prompt. Returns: - List[int]: The suffix of the prompt that actually needs to be processed - by the model. This will be the full prompt if the cache is - reset or cannot be effectively used. + List[Any]: The prompt cache object + List[int]: The tokens that are in the returned object + List[int]: The remaining tokens to be added """ - cache_len = len(self.prompt_cache.tokens) - prompt_len = len(prompt) - com_prefix_len = common_prefix_len(self.prompt_cache.tokens, prompt) - - # Leave at least one token in the prompt - com_prefix_len = min(com_prefix_len, len(prompt) - 1) - - # Condition 1: Model changed or no common prefix at all. Reset cache. - if ( - self.prompt_cache.model_key != self.model_provider.model_key - or com_prefix_len == 0 - ): - self.reset_prompt_cache(prompt) - - # Condition 2: Common prefix exists and matches cache length. Process suffix. - elif com_prefix_len == cache_len: - logging.debug( - f"*** Cache is prefix of prompt (cache_len: {cache_len}, prompt_len: {prompt_len}). Processing suffix. ***" - ) - prompt = prompt[com_prefix_len:] - self.prompt_cache.tokens.extend(prompt) - - # Condition 3: Common prefix exists but is shorter than cache length. Attempt trim. - elif com_prefix_len < cache_len: - logging.debug( - f"*** Common prefix ({com_prefix_len}) shorter than cache ({cache_len}). Attempting trim. ***" - ) - - if can_trim_prompt_cache(self.prompt_cache.cache): - num_to_trim = cache_len - com_prefix_len - logging.debug(f" Trimming {num_to_trim} tokens from cache.") - trim_prompt_cache(self.prompt_cache.cache, num_to_trim) - self.prompt_cache.tokens = self.prompt_cache.tokens[:com_prefix_len] - prompt = prompt[com_prefix_len:] - self.prompt_cache.tokens.extend(prompt) - else: - logging.debug(f" Cache cannot be trimmed. Resetting cache.") - self.reset_prompt_cache(prompt) + cache, rest = self.prompt_cache.fetch_nearest_cache( + self.model_provider.model_key, prompt + ) + cache_key = prompt[: len(prompt) - len(rest)] - # This case should logically not be reached if com_prefix_len <= cache_len - else: - logging.error( - f"Unexpected cache state: com_prefix_len ({com_prefix_len}) > cache_len ({cache_len}). Resetting cache." - ) - self.reset_prompt_cache(prompt) + # Make a new cache for the model + if cache is None: + cache = make_prompt_cache(self.model_provider.model) + if self.model_provider.draft_model is not None: + cache += make_prompt_cache(self.model_provider.draft_model) - logging.debug(f"Returning {len(prompt)} tokens for processing.") - return prompt + return cache, cache_key, rest def handle_completion( self, @@ -645,7 +729,7 @@ def handle_completion( token_logprobs = [] top_tokens = [] - prompt = self.get_prompt_cache(prompt) + cache, cache_key, prompt = self.get_prompt_cache(prompt) text = "" tic = time.perf_counter() @@ -688,6 +772,8 @@ def keepalive_callback(processed_tokens, total_tokens): # Client disconnected, ignore pass + cache_key += prompt + prompt_token_count = len(cache_key) for gen_response in stream_generate( model=self.model, tokenizer=self.tokenizer, @@ -695,7 +781,7 @@ def keepalive_callback(processed_tokens, total_tokens): max_tokens=self.max_tokens, sampler=sampler, logits_processors=logits_processors, - prompt_cache=self.prompt_cache.cache, + prompt_cache=cache, draft_model=self.model_provider.draft_model, num_draft_tokens=self.num_draft_tokens, prompt_progress_callback=keepalive_callback, @@ -720,7 +806,7 @@ def keepalive_callback(processed_tokens, total_tokens): token = gen_response.token logprobs = gen_response.logprobs tokens.append(token) - self.prompt_cache.tokens.append(token) + cache_key.append(token) if self.logprobs > 0: sorted_indices = mx.argpartition(-logprobs, kth=self.logprobs - 1) @@ -777,11 +863,8 @@ def keepalive_callback(processed_tokens, total_tokens): self.wfile.write(f"data: {json.dumps(response)}\n\n".encode()) self.wfile.flush() if self.stream_options is not None and self.stream_options["include_usage"]: - original_prompt_length = ( - len(self.prompt_cache.tokens) - len(tokens) + len(prompt) - ) response = self.completion_usage_response( - original_prompt_length, len(tokens) + prompt_token_count, len(tokens) ) self.wfile.write(f"data: {json.dumps(response)}\n\n".encode()) self.wfile.flush() @@ -791,7 +874,7 @@ def keepalive_callback(processed_tokens, total_tokens): response = self.generate_response( text, finish_reason, - len(prompt), + prompt_token_count, len(tokens), token_logprobs=token_logprobs, top_tokens=top_tokens, @@ -808,6 +891,8 @@ def keepalive_callback(processed_tokens, total_tokens): self.wfile.write(response_json) self.wfile.flush() + self.prompt_cache.insert_cache(self.model_provider.model_key, cache_key, cache) + def completion_usage_response( self, prompt_token_count: Optional[int] = None, @@ -947,7 +1032,7 @@ def run( handler_class=APIHandler, ): server_address = (host, port) - prompt_cache = PromptCache() + prompt_cache = LRUPromptCache() infos = socket.getaddrinfo( *server_address, type=socket.SOCK_STREAM, flags=socket.AI_PASSIVE ) diff --git a/tests/test_server.py b/tests/test_server.py index b28669703..0c819b335 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -6,9 +6,11 @@ import threading import unittest +import mlx.core as mx import requests -from mlx_lm.server import APIHandler +from mlx_lm.models.cache import KVCache +from mlx_lm.server import APIHandler, LRUPromptCache from mlx_lm.utils import load @@ -339,182 +341,6 @@ def test_prompt_cache_with_draft_model(self): self.assertIsNotNone(second_response_body["choices"][0]["message"]["content"]) -# --- Tests for get_prompt_cache --- - -from unittest.mock import MagicMock, patch - -from mlx_lm.server import PromptCache - - -class TestGetPromptCache(unittest.TestCase): - - def setUp(self): - """Set up mocks and a handler instance for each test.""" - self.mock_model_provider = MagicMock() - # Simulate tokenizer needed for decoding in original debug logs (though not strictly needed for cache logic) - self.mock_model_provider.tokenizer = MagicMock() - self.mock_model_provider.tokenizer.decode = lambda x: f"decoded({x})" - self.mock_model_provider.model_key = ("model_v1", None, None) - self.mock_model_provider.draft_model = None # Start without draft model - - # --- Prevent BaseHTTPRequestHandler.__init__ from running --- - # It tries to handle a request immediately, which fails with mocks. - # We only need the APIHandler instance with its attributes set. - with patch( - "http.server.BaseHTTPRequestHandler.__init__", lambda *args, **kwargs: None - ): - # APIHandler init still requires args for BaseHTTPRequestHandler signature, - # but they won't be used by the patched __init__. - mock_request = MagicMock() - mock_client_address = ("127.0.0.1", 8080) - mock_server = MagicMock() - - self.prompt_cache_instance = PromptCache() - self.handler = APIHandler( - self.mock_model_provider, - mock_request, - mock_client_address, - mock_server, - prompt_cache=self.prompt_cache_instance, # Inject our cache instance - ) - # Manually set attributes usually set by APIHandler.__init__ if needed - # self.handler.created = MagicMock() - # self.handler.system_fingerprint = MagicMock() - # (Not strictly necessary for get_prompt_cache testing) - - @patch("mlx_lm.server.make_prompt_cache") - def test_initial_request_empty_cache(self, mock_make_cache): - """Test first request when the cache is empty.""" - mock_make_cache.return_value = "new_cache_obj" - prompt = [1, 2, 3] - - processed_prompt = self.handler.get_prompt_cache(prompt) - - self.assertEqual(processed_prompt, [1, 2, 3]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3]) - self.assertEqual(self.handler.prompt_cache.cache, "new_cache_obj") - self.assertEqual(self.handler.prompt_cache.model_key, ("model_v1", None, None)) - mock_make_cache.assert_called_once() - - @patch("mlx_lm.server.trim_prompt_cache") - @patch("mlx_lm.server.can_trim_prompt_cache", return_value=True) - def test_identical_request_full_hit(self, mock_can_trim, mock_trim_cache): - """Test when the new prompt is identical to the cached one.""" - self.handler.prompt_cache.tokens = [1, 2, 3] - self.handler.prompt_cache.model_key = ("model_v1", None, None) - self.handler.prompt_cache.cache = "existing_cache_obj" - prompt = [1, 2, 3] - - # Mock common_prefix_len to return the full length - with patch("mlx_lm.server.common_prefix_len", return_value=3): - processed_prompt = self.handler.get_prompt_cache(prompt) - - mock_trim_cache.assert_called_once_with("existing_cache_obj", 1) - self.assertEqual(processed_prompt, [3]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3]) - - def test_cache_is_prefix(self): - """Test when the cached prompt is a prefix of the new prompt.""" - self.handler.prompt_cache.tokens = [1, 2, 3] - self.handler.prompt_cache.model_key = ("model_v1", None, None) - self.handler.prompt_cache.cache = "existing_cache_obj" - prompt = [1, 2, 3, 4, 5] - - with patch("mlx_lm.server.common_prefix_len", return_value=3): - processed_prompt = self.handler.get_prompt_cache(prompt) - - # Should process the suffix, cache tokens updated - self.assertEqual(processed_prompt, [4, 5]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3, 4, 5]) - self.assertEqual(self.handler.prompt_cache.cache, "existing_cache_obj") - - @patch("mlx_lm.server.trim_prompt_cache") - @patch("mlx_lm.server.can_trim_prompt_cache", return_value=True) - def test_partial_match_trim_success(self, mock_can_trim, mock_trim_cache): - """Test partial match where cache is longer and trimming succeeds.""" - self.handler.prompt_cache.tokens = [1, 2, 3, 4, 5] - self.handler.prompt_cache.model_key = ("model_v1", None, None) - self.handler.prompt_cache.cache = "existing_cache_obj" - prompt = [1, 2, 3, 6, 7] # Diverges after token 3 - - with patch("mlx_lm.server.common_prefix_len", return_value=3): - processed_prompt = self.handler.get_prompt_cache(prompt) - - # Should process the new suffix, cache trimmed and updated - self.assertEqual(processed_prompt, [6, 7]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3, 6, 7]) - mock_can_trim.assert_called_once_with("existing_cache_obj") - # Called with cache object and num_to_trim (5 - 3 = 2) - mock_trim_cache.assert_called_once_with("existing_cache_obj", 2) - self.assertEqual( - self.handler.prompt_cache.cache, "existing_cache_obj" - ) # Cache obj itself isn't changed by mock - - @patch("mlx_lm.server.make_prompt_cache") - @patch("mlx_lm.server.trim_prompt_cache") - @patch("mlx_lm.server.can_trim_prompt_cache", return_value=False) - def test_partial_match_trim_fail( - self, mock_can_trim, mock_trim_cache, mock_make_cache - ): - """Test partial match where cache is longer but trimming fails.""" - mock_make_cache.return_value = "new_cache_obj_on_reset" - self.handler.prompt_cache.tokens = [1, 2, 3, 4, 5] - self.handler.prompt_cache.model_key = ("model_v1", None, None) - self.handler.prompt_cache.cache = "existing_cache_obj" - prompt = [1, 2, 3, 6, 7] # Diverges after token 3 - - with patch("mlx_lm.server.common_prefix_len", return_value=3): - processed_prompt = self.handler.get_prompt_cache(prompt) - - # Should process the full prompt, cache reset - self.assertEqual(processed_prompt, [1, 2, 3, 6, 7]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3, 6, 7]) - mock_can_trim.assert_called_once_with("existing_cache_obj") - mock_trim_cache.assert_not_called() - mock_make_cache.assert_called_once() # Cache was reset - self.assertEqual(self.handler.prompt_cache.cache, "new_cache_obj_on_reset") - - @patch("mlx_lm.server.make_prompt_cache") - def test_no_common_prefix(self, mock_make_cache): - """Test when there is no common prefix between cache and prompt.""" - mock_make_cache.return_value = "new_cache_obj" - self.handler.prompt_cache.tokens = [1, 2, 3] - self.handler.prompt_cache.model_key = ("model_v1", None, None) - self.handler.prompt_cache.cache = "existing_cache_obj" - prompt = [4, 5, 6] - - with patch("mlx_lm.server.common_prefix_len", return_value=0): - processed_prompt = self.handler.get_prompt_cache(prompt) - - # Should process the full prompt, cache reset - self.assertEqual(processed_prompt, [4, 5, 6]) - self.assertEqual(self.handler.prompt_cache.tokens, [4, 5, 6]) - mock_make_cache.assert_called_once() - self.assertEqual(self.handler.prompt_cache.cache, "new_cache_obj") - - @patch("mlx_lm.server.make_prompt_cache") - def test_model_changed(self, mock_make_cache): - """Test cache reset when the model key changes.""" - mock_make_cache.return_value = "new_cache_obj_model_change" - self.handler.prompt_cache.tokens = [1, 2, 3] - self.handler.prompt_cache.model_key = ("model_v1", None, None) # Original key - self.handler.prompt_cache.cache = "existing_cache_obj" - - # Simulate model provider having a new key - self.mock_model_provider.model_key = ("model_v2", None, None) - prompt = [1, 2, 3, 4] - - # No need to mock common_prefix_len, model check happens first - processed_prompt = self.handler.get_prompt_cache(prompt) - - # Should process the full prompt, cache reset - self.assertEqual(processed_prompt, [1, 2, 3, 4]) - self.assertEqual(self.handler.prompt_cache.tokens, [1, 2, 3, 4]) - mock_make_cache.assert_called_once() - self.assertEqual(self.handler.prompt_cache.cache, "new_cache_obj_model_change") - self.assertEqual(self.handler.prompt_cache.model_key, ("model_v2", None, None)) - - class TestKeepalive(unittest.TestCase): def test_keepalive_callback(self): @@ -565,5 +391,78 @@ def keepalive_callback(processed_tokens, total_tokens): self.fail(f"Callback should handle BrokenPipeError: {e}") +class TestLRUPromptCache(unittest.TestCase): + + def test_caching(self): + cache = LRUPromptCache(max_size=10) + + def get_kv(n): + keys = mx.arange(n).reshape(1, 1, n, 1) + return keys, keys + + model = ("test", None, None) + tokens = [10] * 24 + + c, t = cache.fetch_nearest_cache(model, tokens) + self.assertTrue(c is None) + self.assertEqual(t, tokens) + + c = [KVCache()] + c[0].update_and_fetch(*get_kv(24)) + cache.insert_cache(model, t, c) + + tokens = tokens + [20] * 5 + c, t = cache.fetch_nearest_cache(model, tokens) + k, v = c[0].state + self.assertTrue((k == v).all().item()) + self.assertTrue((k.flatten() == mx.arange(24)).all().item()) + self.assertEqual(t, [20] * 5) + self.assertEqual(len(cache._lru), 0) + + tokens = tokens + [30] * 3 + c[0].update_and_fetch(*get_kv(8)) + cache.insert_cache(model, tokens, c) + + tokens = tokens[:26] + [40] * 8 + c, t = cache.fetch_nearest_cache(model, tokens) + k, v = c[0].state + self.assertTrue((k == v).all().item()) + self.assertTrue( + (k.flatten() == mx.concatenate([mx.arange(24), mx.arange(2)])).all().item() + ) + self.assertEqual(t, [40] * 8) + self.assertEqual(len(cache._lru), 1) + + def test_lru(self): + cache = LRUPromptCache(max_size=2) + model = ("test", None, None) + cache.insert_cache(model, [1, 2], ["test1"]) + cache.insert_cache(model, [1, 2], ["test1"]) + + c, t = cache.fetch_nearest_cache(model, [1, 2]) + self.assertEqual(c, ["test1"]) + self.assertEqual(t, []) + c, t = cache.fetch_nearest_cache(model, [1, 2]) + self.assertEqual(c, ["test1"]) + self.assertEqual(t, []) + c, t = cache.fetch_nearest_cache(model, [1, 2]) + self.assertEqual(c, None) + self.assertEqual(t, [1, 2]) + + cache.insert_cache(model, [1, 2], ["test1"]) + cache.insert_cache(model, [2, 3], ["test2"]) + cache.insert_cache(model, [3, 4], ["test3"]) + + c, t = cache.fetch_nearest_cache(model, [1, 2]) + self.assertEqual(c, None) + self.assertEqual(t, [1, 2]) + c, t = cache.fetch_nearest_cache(model, [2, 3]) + self.assertEqual(c, ["test2"]) + self.assertEqual(t, []) + c, t = cache.fetch_nearest_cache(model, [3, 4]) + self.assertEqual(c, ["test3"]) + self.assertEqual(t, []) + + if __name__ == "__main__": unittest.main()