Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
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
25 changes: 19 additions & 6 deletions mlx_lm/models/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -1607,15 +1607,28 @@ def search(self, model: Any, tokens: List[int]) -> PromptTrieResult:
common_prefix = index
if index > 0:
best = None
stack = [(current, [])]
# Reuse one path while traversing instead of copying the growing
# path for every node. Markers restore it after each child.
path = []
pop_path = object()
stack = [(current, None, False)]
while stack:
current, extra = stack.pop()
item = stack.pop()
if item is pop_path:
path.pop()
continue

current, tok, append_token = item
if append_token:
path.append(tok)

if "__value__" in current:
if best is None or len(extra) < len(best):
best = extra
elif best is None or len(extra) < len(best):
if best is None or len(path) < len(best):
best = list(path)
elif best is None or len(path) < len(best):
for tok in current:
stack.append((current[tok], extra + [tok]))
stack.append(pop_path)
stack.append((current[tok], tok, True))
longer = tokens[:index] + best
return PromptTrieResult(model, None, shorter, longer, common_prefix)

Expand Down
25 changes: 24 additions & 1 deletion tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
import requests

from mlx_lm.generate import TextStateMachine
from mlx_lm.models.cache import KVCache
from mlx_lm.models.cache import KVCache, PromptTrie
from mlx_lm.server import (
APIHandler,
LRUPromptCache,
Expand Down Expand Up @@ -565,6 +565,29 @@ def keepalive_callback(processed_tokens, total_tokens):
self.fail(f"Callback should handle BrokenPipeError: {e}")


class TestPromptTrie(unittest.TestCase):
def test_search_finds_shortest_longer_sequence(self):
trie = PromptTrie()
model = object()
trie.add(model, [1, 2, 3, 4], "long")
trie.add(model, [1, 5, 6], "short")

result = trie.search(model, [1])

self.assertEqual(result.longer, [1, 5, 6])
self.assertEqual(result.common_prefix, 1)

def test_search_longer_sequence_tie_breaking(self):
trie = PromptTrie()
model = object()
trie.add(model, [1, 2, 3], "first")
trie.add(model, [1, 4, 5], "second")

result = trie.search(model, [1])

self.assertEqual(result.longer, [1, 4, 5])


class TestLRUPromptCache(unittest.TestCase):
def test_caching(self):
cache = LRUPromptCache(max_size=10)
Expand Down