Studio: reuse MLX prompt cache across turns instead of re-prefilling (#7311)
* Studio: reuse MLX prompt cache across turns instead of re-prefilling * clean up * key prompt cache on what the KV covers * skip windowed KV caches past their window * verify prefix coverage before caching KV
This commit is contained in:
parent
8517721adb
commit
59bda2e1f7
2 changed files with 611 additions and 13 deletions
|
|
@ -181,19 +181,27 @@ def _vlm_messages_have_tool_history(messages):
|
|||
)
|
||||
|
||||
|
||||
def _build_generation_stats(prompt_n, prompt_tps, gen_n, gen_tps):
|
||||
def _build_generation_stats(
|
||||
prompt_n,
|
||||
prompt_tps,
|
||||
gen_n,
|
||||
gen_tps,
|
||||
cached_n = 0,
|
||||
):
|
||||
"""Map mlx stream stats onto the usage/timings shape llama-server emits."""
|
||||
prompt_n = int(prompt_n or 0)
|
||||
gen_n = int(gen_n or 0)
|
||||
cached_n = int(cached_n or 0)
|
||||
prompt_tps = float(prompt_tps or 0.0)
|
||||
gen_tps = float(gen_tps or 0.0)
|
||||
prompt_ms = (prompt_n / prompt_tps * 1000.0) if prompt_tps > 0 else 0.0
|
||||
predicted_ms = (gen_n / gen_tps * 1000.0) if gen_tps > 0 else 0.0
|
||||
total_prompt_n = prompt_n + cached_n
|
||||
return {
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_n,
|
||||
"prompt_tokens": total_prompt_n,
|
||||
"completion_tokens": gen_n,
|
||||
"total_tokens": prompt_n + gen_n,
|
||||
"total_tokens": total_prompt_n + gen_n,
|
||||
},
|
||||
"timings": {
|
||||
"prompt_n": prompt_n,
|
||||
|
|
@ -204,11 +212,123 @@ def _build_generation_stats(prompt_n, prompt_tps, gen_n, gen_tps):
|
|||
"predicted_ms": predicted_ms,
|
||||
"predicted_per_token_ms": (predicted_ms / gen_n) if gen_n > 0 else 0.0,
|
||||
"predicted_per_second": gen_tps,
|
||||
"cache_n": 0,
|
||||
"cache_n": cached_n,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
PROMPT_CACHE_ENTRIES = 6
|
||||
PROMPT_CACHE_MEMORY_FRACTION = 0.15
|
||||
PROMPT_CACHE_FALLBACK_BYTES = 2 * 1024**3
|
||||
|
||||
|
||||
def _mlx_prompt_cache_api():
|
||||
try:
|
||||
from mlx_lm.models.cache import (
|
||||
LRUPromptCache,
|
||||
can_trim_prompt_cache,
|
||||
make_prompt_cache,
|
||||
trim_prompt_cache,
|
||||
)
|
||||
except ImportError:
|
||||
return None
|
||||
return LRUPromptCache, make_prompt_cache, can_trim_prompt_cache, trim_prompt_cache
|
||||
|
||||
|
||||
def _prompt_cache_max_bytes(recommended_gb = None):
|
||||
override = os.environ.get("UNSLOTH_MLX_PROMPT_CACHE_BYTES")
|
||||
if override:
|
||||
try:
|
||||
return max(int(override), 0)
|
||||
except ValueError:
|
||||
logger.warning("Ignoring non-integer UNSLOTH_MLX_PROMPT_CACHE_BYTES=%r", override)
|
||||
if recommended_gb:
|
||||
return int(recommended_gb * 1e9 * PROMPT_CACHE_MEMORY_FRACTION)
|
||||
return PROMPT_CACHE_FALLBACK_BYTES
|
||||
|
||||
|
||||
def _flatten_kv_entries(cache):
|
||||
for entry in cache:
|
||||
nested = getattr(entry, "caches", None)
|
||||
if nested is None:
|
||||
yield entry
|
||||
else:
|
||||
yield from _flatten_kv_entries(nested)
|
||||
|
||||
|
||||
def _kv_prefix_coverage(cache):
|
||||
covered = None
|
||||
for entry in _flatten_kv_entries(cache):
|
||||
offset = getattr(entry, "offset", None)
|
||||
if offset is None:
|
||||
return None
|
||||
if getattr(entry, "start_position", 0):
|
||||
return None
|
||||
window = getattr(entry, "max_size", None)
|
||||
if window is not None and offset > window:
|
||||
return None
|
||||
if covered is None:
|
||||
covered = offset
|
||||
elif covered != offset:
|
||||
return None
|
||||
return covered
|
||||
|
||||
|
||||
class _MLXPromptCacheHistory:
|
||||
def __init__(self, max_entries, max_bytes):
|
||||
api = _mlx_prompt_cache_api()
|
||||
if api is None:
|
||||
raise RuntimeError("mlx-lm is too old for LRUPromptCache")
|
||||
lru_cls, make, can_trim, trim = api
|
||||
self._make_prompt_cache = make
|
||||
self._can_trim = can_trim
|
||||
self._trim = trim
|
||||
self._max_bytes = max_bytes
|
||||
self._lru = lru_cls(max_size = max_entries, max_bytes = max_bytes)
|
||||
|
||||
def fetch(self, model, key, tokens):
|
||||
cache, rest = self._lru.fetch_nearest_cache(key, list(tokens))
|
||||
if cache is not None:
|
||||
if rest:
|
||||
return cache, list(rest)
|
||||
if self._can_trim(cache) and self._trim(cache, 1) == 1:
|
||||
return cache, list(tokens[-1:])
|
||||
if len(tokens) > 1:
|
||||
head = list(tokens[:-1])
|
||||
cache, rest = self._lru.fetch_nearest_cache(key, head)
|
||||
if cache is not None:
|
||||
covered = len(head) - len(rest)
|
||||
return cache, list(tokens[covered:])
|
||||
return self._make_prompt_cache(model), list(tokens)
|
||||
|
||||
def insert(self, key, tokens, cache):
|
||||
# An over-budget entry evicts itself and every other conversation.
|
||||
nbytes = sum(getattr(entry, "nbytes", 0) for entry in cache)
|
||||
if nbytes > self._max_bytes:
|
||||
logger.debug(
|
||||
"MLX prompt cache: skipping %.2f GB entry over the %.2f GB budget",
|
||||
nbytes / 1e9,
|
||||
self._max_bytes / 1e9,
|
||||
)
|
||||
return
|
||||
covered = _kv_prefix_coverage(cache)
|
||||
if covered is None:
|
||||
logger.debug("MLX prompt cache: skipping cache with unverifiable prefix coverage")
|
||||
return
|
||||
tokens = list(tokens)
|
||||
if covered > len(tokens):
|
||||
logger.debug(
|
||||
"MLX prompt cache: cache covers %d tokens but only %d were tracked",
|
||||
covered,
|
||||
len(tokens),
|
||||
)
|
||||
return
|
||||
tokens = tokens[:covered]
|
||||
if not tokens:
|
||||
return
|
||||
self._lru.insert_cache(key, tokens, cache)
|
||||
|
||||
|
||||
def _mlx_distributed_rank_size(group = None):
|
||||
"""Return ``(rank, world_size)`` for an optional MLX distributed group."""
|
||||
if group is None:
|
||||
|
|
@ -313,6 +433,55 @@ class MLXInferenceBackend:
|
|||
# Recorded for unload to release pinned memory back to the OS.
|
||||
self._memory_limits_applied = {}
|
||||
|
||||
self._prompt_cache_history = None
|
||||
self._prompt_cache_unavailable = False
|
||||
|
||||
def _prompt_cache(self):
|
||||
if self._prompt_cache_history is not None or self._prompt_cache_unavailable:
|
||||
return self._prompt_cache_history
|
||||
max_bytes = _prompt_cache_max_bytes(self._memory_limits_applied.get("recommended_gb"))
|
||||
if max_bytes <= 0:
|
||||
self._prompt_cache_unavailable = True
|
||||
logger.info("MLX prompt cache disabled by budget")
|
||||
return None
|
||||
try:
|
||||
self._prompt_cache_history = _MLXPromptCacheHistory(
|
||||
PROMPT_CACHE_ENTRIES,
|
||||
max_bytes,
|
||||
)
|
||||
except Exception as exc:
|
||||
self._prompt_cache_unavailable = True
|
||||
logger.info("MLX prompt cache unavailable (%s); prefilling every request", exc)
|
||||
return None
|
||||
logger.info(
|
||||
"MLX prompt cache: %d entries, %.2f GB budget",
|
||||
PROMPT_CACHE_ENTRIES,
|
||||
max_bytes / 1e9,
|
||||
)
|
||||
return self._prompt_cache_history
|
||||
|
||||
def _clear_prompt_cache(self):
|
||||
self._prompt_cache_history = None
|
||||
self._prompt_cache_unavailable = False
|
||||
|
||||
def _prepare_prompt_cache(self, prompt, adapter_state):
|
||||
history = self._prompt_cache()
|
||||
if history is None:
|
||||
return prompt, None, None, None, 0
|
||||
try:
|
||||
tokenizer = self._tokenizer
|
||||
bos = getattr(tokenizer, "bos_token", None)
|
||||
add_special_tokens = bos is None or not prompt.startswith(bos)
|
||||
tokens = list(tokenizer.encode(prompt, add_special_tokens = add_special_tokens))
|
||||
if not tokens:
|
||||
return prompt, None, None, None, 0
|
||||
key = f"{self.active_model_name}|{adapter_state!r}"
|
||||
cache, rest = history.fetch(self._model, key, tokens)
|
||||
except Exception as exc:
|
||||
logger.debug("MLX prompt cache lookup failed: %s", exc)
|
||||
return prompt, None, None, None, 0
|
||||
return rest, cache, key, tokens, len(tokens) - len(rest)
|
||||
|
||||
def _configure_memory_limits(self):
|
||||
"""Apply Metal memory caps before loading a model.
|
||||
|
||||
|
|
@ -535,6 +704,7 @@ class MLXInferenceBackend:
|
|||
self._distributed_world_size = 1
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
self._clear_prompt_cache()
|
||||
gc.collect()
|
||||
mx.clear_cache()
|
||||
|
||||
|
|
@ -731,24 +901,34 @@ class MLXInferenceBackend:
|
|||
# <think> prefix on every native-protocol snapshot just as the normal
|
||||
# decoding path does below.
|
||||
normalized_output = think_prefix
|
||||
logger.info(
|
||||
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
|
||||
len(prompt),
|
||||
max_new_tokens,
|
||||
type(self._model).__name__,
|
||||
type(self._tokenizer).__name__,
|
||||
)
|
||||
with self._generation_lock, _temporary_mlx_adapter_state(self._model, _adapter_state):
|
||||
(
|
||||
gen_prompt,
|
||||
prompt_cache,
|
||||
cache_key,
|
||||
prompt_tokens,
|
||||
cached_n,
|
||||
) = self._prepare_prompt_cache(prompt, _adapter_state)
|
||||
logger.info(
|
||||
"Generating: prompt_len=%d, cached=%d, max_tokens=%d, model=%s, tokenizer=%s",
|
||||
len(prompt),
|
||||
cached_n,
|
||||
max_new_tokens,
|
||||
type(self._model).__name__,
|
||||
type(self._tokenizer).__name__,
|
||||
)
|
||||
final_response = None
|
||||
try:
|
||||
# Enter request-scoped model state before yielding any response.
|
||||
if think_prefix:
|
||||
yield think_prefix
|
||||
gen_kwargs = dict(
|
||||
prompt = prompt,
|
||||
prompt = gen_prompt,
|
||||
max_tokens = max_new_tokens,
|
||||
sampler = sampler,
|
||||
)
|
||||
if prompt_cache is not None:
|
||||
gen_kwargs["prompt_cache"] = prompt_cache
|
||||
if logits_processors is not None:
|
||||
gen_kwargs["logits_processors"] = logits_processors
|
||||
for response in stream_generate(
|
||||
|
|
@ -757,6 +937,7 @@ class MLXInferenceBackend:
|
|||
**gen_kwargs,
|
||||
):
|
||||
final_response = response
|
||||
token_ids.append(response.token)
|
||||
if preserve_native_channels:
|
||||
piece = getattr(response, "text", None) or ""
|
||||
delta = normalizer.feed(piece)
|
||||
|
|
@ -764,7 +945,6 @@ class MLXInferenceBackend:
|
|||
normalized_output += delta
|
||||
yield normalized_output
|
||||
else:
|
||||
token_ids.append(response.token)
|
||||
cumulative = self._tokenizer.decode(
|
||||
token_ids,
|
||||
skip_special_tokens = True,
|
||||
|
|
@ -773,6 +953,13 @@ class MLXInferenceBackend:
|
|||
|
||||
if cancel_event and cancel_event.is_set():
|
||||
break
|
||||
if prompt_cache is not None and prompt_tokens is not None:
|
||||
history = self._prompt_cache_history
|
||||
if history is not None:
|
||||
try:
|
||||
history.insert(cache_key, prompt_tokens + token_ids, prompt_cache)
|
||||
except Exception as exc:
|
||||
logger.debug("MLX prompt cache insert failed: %s", exc)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
logger.error("stream_generate failed:\n%s", traceback.format_exc())
|
||||
|
|
@ -785,6 +972,7 @@ class MLXInferenceBackend:
|
|||
getattr(final_response, "prompt_tps", 0.0),
|
||||
getattr(final_response, "generation_tokens", 0),
|
||||
getattr(final_response, "generation_tps", 0.0),
|
||||
cached_n,
|
||||
)
|
||||
if normalizer is not None:
|
||||
cancelled = cancel_event is not None and cancel_event.is_set()
|
||||
|
|
|
|||
|
|
@ -922,3 +922,413 @@ def test_mlx_vlm_normalizes_native_reasoning_channels(monkeypatch):
|
|||
"<think>vision</think>",
|
||||
"<think>vision</think> answer",
|
||||
]
|
||||
|
||||
|
||||
class _FakeLRUPromptCache:
|
||||
def __init__(
|
||||
self,
|
||||
max_size = 10,
|
||||
max_bytes = 1 << 63,
|
||||
):
|
||||
self.max_size = max_size
|
||||
self.max_bytes = max_bytes
|
||||
self.entries = {}
|
||||
|
||||
def fetch_nearest_cache(self, key, tokens):
|
||||
import copy
|
||||
|
||||
stored = self.entries.get(key, {})
|
||||
exact = stored.get(tuple(tokens))
|
||||
if exact is not None:
|
||||
return copy.deepcopy(exact), []
|
||||
best = None
|
||||
for candidate, cache in stored.items():
|
||||
if len(candidate) < len(tokens) and tuple(tokens[: len(candidate)]) == candidate:
|
||||
if best is None or len(candidate) > len(best[0]):
|
||||
best = (candidate, cache)
|
||||
if best is not None:
|
||||
return copy.deepcopy(best[1]), list(tokens[len(best[0]) :])
|
||||
return None, list(tokens)
|
||||
|
||||
def insert_cache(
|
||||
self,
|
||||
key,
|
||||
tokens,
|
||||
prompt_cache,
|
||||
*,
|
||||
cache_type = "assistant",
|
||||
):
|
||||
import copy
|
||||
self.entries.setdefault(key, {})[tuple(tokens)] = copy.deepcopy(prompt_cache)
|
||||
|
||||
|
||||
class _FakeCacheEntry:
|
||||
def __init__(
|
||||
self,
|
||||
offset = 0,
|
||||
nbytes = 1,
|
||||
):
|
||||
self.offset = offset
|
||||
self.nbytes = nbytes
|
||||
|
||||
|
||||
def _install_fake_prompt_cache_api(monkeypatch, trimmable = True):
|
||||
from core.inference import mlx_inference
|
||||
|
||||
def _make_prompt_cache(_model):
|
||||
return [_FakeCacheEntry()]
|
||||
|
||||
def _can_trim_prompt_cache(_cache):
|
||||
return trimmable
|
||||
|
||||
def _trim_prompt_cache(cache, num):
|
||||
cache[0].offset = max(cache[0].offset - num, 0)
|
||||
return num
|
||||
|
||||
monkeypatch.setattr(
|
||||
mlx_inference,
|
||||
"_mlx_prompt_cache_api",
|
||||
lambda: (
|
||||
_FakeLRUPromptCache,
|
||||
_make_prompt_cache,
|
||||
_can_trim_prompt_cache,
|
||||
_trim_prompt_cache,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_max_bytes_budget(monkeypatch):
|
||||
from core.inference.mlx_inference import (
|
||||
PROMPT_CACHE_FALLBACK_BYTES,
|
||||
PROMPT_CACHE_MEMORY_FRACTION,
|
||||
_prompt_cache_max_bytes,
|
||||
)
|
||||
|
||||
monkeypatch.delenv("UNSLOTH_MLX_PROMPT_CACHE_BYTES", raising = False)
|
||||
assert _prompt_cache_max_bytes(None) == PROMPT_CACHE_FALLBACK_BYTES
|
||||
assert _prompt_cache_max_bytes(20.0) == int(20.0 * 1e9 * PROMPT_CACHE_MEMORY_FRACTION)
|
||||
|
||||
monkeypatch.setenv("UNSLOTH_MLX_PROMPT_CACHE_BYTES", "4096")
|
||||
assert _prompt_cache_max_bytes(20.0) == 4096
|
||||
monkeypatch.setenv("UNSLOTH_MLX_PROMPT_CACHE_BYTES", "0")
|
||||
assert _prompt_cache_max_bytes(20.0) == 0
|
||||
monkeypatch.setenv("UNSLOTH_MLX_PROMPT_CACHE_BYTES", "not-a-number")
|
||||
assert _prompt_cache_max_bytes(20.0) == int(20.0 * 1e9 * PROMPT_CACHE_MEMORY_FRACTION)
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_never_returns_empty_remainder(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
from core.inference.mlx_inference import _MLXPromptCacheHistory
|
||||
|
||||
history = _MLXPromptCacheHistory(6, 1 << 30)
|
||||
tokens = list(range(10))
|
||||
cache, rest = history.fetch(object(), "key", tokens)
|
||||
assert len(rest) == 10
|
||||
cache[0].offset = len(tokens)
|
||||
history.insert("key", tokens, cache)
|
||||
|
||||
_cache, rest = history.fetch(object(), "key", tokens)
|
||||
assert rest == tokens[-1:]
|
||||
|
||||
longer = tokens + [99, 100]
|
||||
_cache, rest = history.fetch(object(), "key", longer)
|
||||
assert rest == [99, 100]
|
||||
|
||||
_install_fake_prompt_cache_api(monkeypatch, trimmable = False)
|
||||
history = _MLXPromptCacheHistory(6, 1 << 30)
|
||||
cache, _rest = history.fetch(object(), "key", tokens)
|
||||
cache[0].offset = len(tokens)
|
||||
history.insert("key", tokens, cache)
|
||||
_cache, rest = history.fetch(object(), "key", tokens)
|
||||
assert rest == tokens, "untrimmable entry must not be reused"
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_key_isolates_adapter_state(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
_install_fake_mlx(monkeypatch)
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
class _Tok:
|
||||
bos_token = None
|
||||
|
||||
def encode(
|
||||
self,
|
||||
text,
|
||||
add_special_tokens = True,
|
||||
):
|
||||
return [ord(c) for c in text]
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
backend._model = object()
|
||||
backend._tokenizer = _Tok()
|
||||
backend.active_model_name = "model-a"
|
||||
|
||||
prompt = "shared prefix"
|
||||
_rest, cache, key, tokens, cached = backend._prepare_prompt_cache(prompt, True)
|
||||
assert cached == 0
|
||||
cache[0].offset = len(tokens)
|
||||
backend._prompt_cache_history.insert(key, tokens, cache)
|
||||
|
||||
_rest, _cache, _key, _tokens, cached_same = backend._prepare_prompt_cache(prompt, True)
|
||||
assert cached_same > 0
|
||||
_rest, _cache, _key, _tokens, cached_flipped = backend._prepare_prompt_cache(prompt, False)
|
||||
assert cached_flipped == 0
|
||||
|
||||
|
||||
def _install_fake_text_stack(
|
||||
monkeypatch,
|
||||
token_map,
|
||||
captured,
|
||||
markers = None,
|
||||
):
|
||||
import types as _types
|
||||
|
||||
from core.inference import mlx_inference
|
||||
|
||||
_install_fake_mlx(monkeypatch)
|
||||
monkeypatch.setattr(
|
||||
mlx_inference,
|
||||
"_temporary_mlx_adapter_state",
|
||||
lambda _model, _state: __import__("contextlib").nullcontext(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"core.inference.chat_template_helpers.apply_chat_template_for_generation",
|
||||
lambda _tok, messages, **_kw: messages[-1]["content"],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"core.inference.chat_template_helpers.render_with_native_template_fallback",
|
||||
lambda formatted_prompt, **_kw: SimpleNamespace(
|
||||
prompt = formatted_prompt,
|
||||
reasoning_channel_markers = markers,
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"core.inference.chat_template_helpers.detect_think_prefill",
|
||||
lambda *_a, **_kw: "",
|
||||
)
|
||||
|
||||
class _Resp:
|
||||
def __init__(self, token, processed):
|
||||
self.token = token
|
||||
self.text = f"<{token}>"
|
||||
self.prompt_tokens = processed
|
||||
self.prompt_tps = 10.0
|
||||
self.generation_tokens = 1
|
||||
self.generation_tps = 5.0
|
||||
|
||||
def _stream_generate(_model, _tokenizer, **kwargs):
|
||||
captured.append(kwargs)
|
||||
processed = len(kwargs["prompt"])
|
||||
cache = kwargs.get("prompt_cache")
|
||||
if cache is not None:
|
||||
cache[0].offset += processed
|
||||
for token in token_map["generated"]:
|
||||
if cache is not None:
|
||||
cache[0].offset += 1
|
||||
yield _Resp(token, processed)
|
||||
|
||||
mlx_lm_pkg = _types.ModuleType("mlx_lm")
|
||||
mlx_lm_pkg.stream_generate = _stream_generate
|
||||
mlx_lm_sample = _types.ModuleType("mlx_lm.sample_utils")
|
||||
mlx_lm_sample.make_sampler = lambda **_kw: object()
|
||||
mlx_lm_sample.make_logits_processors = lambda **_kw: []
|
||||
monkeypatch.setitem(sys.modules, "mlx_lm", mlx_lm_pkg)
|
||||
monkeypatch.setitem(sys.modules, "mlx_lm.sample_utils", mlx_lm_sample)
|
||||
|
||||
class _Tok:
|
||||
bos_token = None
|
||||
chat_template = "x"
|
||||
|
||||
def encode(
|
||||
self,
|
||||
text,
|
||||
add_special_tokens = True,
|
||||
):
|
||||
return list(token_map[text])
|
||||
|
||||
def decode(
|
||||
self,
|
||||
ids,
|
||||
skip_special_tokens = False,
|
||||
):
|
||||
return "".join(str(i) for i in ids)
|
||||
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
backend._model = object()
|
||||
backend._tokenizer = _Tok()
|
||||
backend._is_vlm = False
|
||||
backend.active_model_name = "model-a"
|
||||
return backend
|
||||
|
||||
|
||||
def _run_turn(backend, prompt):
|
||||
list(
|
||||
backend.generate_chat_response(
|
||||
messages = [{"role": "user", "content": prompt}],
|
||||
max_new_tokens = 4,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_mlx_text_reuses_prompt_cache_on_the_next_turn(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
captured = []
|
||||
token_map = {
|
||||
"P1": [1, 2, 3],
|
||||
"P2": [1, 2, 3, 7, 8, 9, 10],
|
||||
"generated": [7, 8],
|
||||
}
|
||||
backend = _install_fake_text_stack(monkeypatch, token_map, captured)
|
||||
|
||||
_run_turn(backend, "P1")
|
||||
assert captured[0]["prompt"] == [1, 2, 3]
|
||||
assert "prompt_cache" in captured[0]
|
||||
assert backend.last_generation_stats["timings"]["cache_n"] == 0
|
||||
|
||||
_run_turn(backend, "P2")
|
||||
assert captured[1]["prompt"] == [9, 10], "turn two should prefill only the new tail"
|
||||
|
||||
stats = backend.last_generation_stats
|
||||
assert stats["timings"]["cache_n"] == 5
|
||||
assert stats["timings"]["prompt_n"] == 2
|
||||
assert stats["usage"]["prompt_tokens"] == 7
|
||||
|
||||
|
||||
def test_mlx_text_without_lru_prompt_cache_prefills_the_full_prompt(monkeypatch):
|
||||
from core.inference import mlx_inference
|
||||
|
||||
monkeypatch.setattr(mlx_inference, "_mlx_prompt_cache_api", lambda: None)
|
||||
captured = []
|
||||
token_map = {"P1": [1, 2, 3], "generated": [7]}
|
||||
backend = _install_fake_text_stack(monkeypatch, token_map, captured)
|
||||
|
||||
_run_turn(backend, "P1")
|
||||
assert captured[0]["prompt"] == "P1"
|
||||
assert "prompt_cache" not in captured[0]
|
||||
assert backend.last_generation_stats["timings"]["cache_n"] == 0
|
||||
|
||||
|
||||
def test_mlx_text_tracks_tokens_on_the_native_reasoning_path(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
captured = []
|
||||
token_map = {"P1": [1, 2, 3], "P2": [1, 2, 3, 7, 8, 9], "generated": [7, 8]}
|
||||
backend = _install_fake_text_stack(monkeypatch, token_map, captured, markers = ("<a>", "</a>"))
|
||||
|
||||
_run_turn(backend, "P1")
|
||||
_run_turn(backend, "P2")
|
||||
assert captured[1]["prompt"] == [9]
|
||||
|
||||
|
||||
def test_mlx_presence_penalty_latches_the_first_decode_step():
|
||||
mx = pytest.importorskip("mlx.core")
|
||||
import numpy as np
|
||||
|
||||
from core.inference.mlx_inference import _make_mlx_presence_penalty_processor
|
||||
|
||||
processor = _make_mlx_presence_penalty_processor(2.0)
|
||||
logits = mx.zeros((1, 5))
|
||||
out = processor(mx.array([3]), logits)
|
||||
assert np.array_equal(np.array(out), np.zeros((1, 5))), "prompt must not be penalized"
|
||||
out = processor(mx.array([3, 1]), mx.zeros((1, 5)))
|
||||
penalized = np.array(out)[0]
|
||||
assert penalized[1] == -2.0
|
||||
assert penalized[3] == 0.0
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_survives_reset_but_not_unload(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
_install_fake_mlx(monkeypatch)
|
||||
sys.modules["mlx.core"].clear_cache = lambda: None
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
backend.active_model_name = "model-a"
|
||||
history = backend._prompt_cache()
|
||||
assert history is not None
|
||||
|
||||
backend.reset_generation_state()
|
||||
assert backend._prompt_cache_history is history
|
||||
|
||||
backend.unload_model("model-a")
|
||||
assert backend._prompt_cache_history is None
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_skips_entries_over_budget(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
from core.inference.mlx_inference import _MLXPromptCacheHistory
|
||||
|
||||
history = _MLXPromptCacheHistory(6, 1000)
|
||||
history.insert("key", [1, 2, 3], [_FakeCacheEntry(offset = 3, nbytes = 400)])
|
||||
assert len(history._lru.entries.get("key", {})) == 1
|
||||
|
||||
history.insert("key", list(range(50)), [_FakeCacheEntry(offset = 50, nbytes = 5000)])
|
||||
stored = history._lru.entries.get("key", {})
|
||||
assert tuple([1, 2, 3]) in stored
|
||||
assert tuple(range(50)) not in stored
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_keys_on_what_the_kv_covers(monkeypatch):
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
from core.inference.mlx_inference import _MLXPromptCacheHistory
|
||||
|
||||
class _Entry:
|
||||
def __init__(
|
||||
self,
|
||||
offset,
|
||||
nbytes = 1,
|
||||
):
|
||||
self.offset = offset
|
||||
self.nbytes = nbytes
|
||||
|
||||
history = _MLXPromptCacheHistory(6, 1 << 30)
|
||||
|
||||
history.insert("key", list(range(10)), [_Entry(offset = 8)])
|
||||
assert tuple(range(8)) in history._lru.entries["key"]
|
||||
assert tuple(range(10)) not in history._lru.entries["key"]
|
||||
|
||||
history.insert("other", list(range(4)), [_Entry(offset = 9)])
|
||||
assert "other" not in history._lru.entries
|
||||
|
||||
|
||||
def test_mlx_prompt_cache_only_stores_verifiable_prefix_coverage(monkeypatch):
|
||||
mx = pytest.importorskip("mlx.core")
|
||||
from mlx_lm.models.cache import CacheList, ChunkedKVCache, KVCache, RotatingKVCache
|
||||
|
||||
_install_fake_prompt_cache_api(monkeypatch)
|
||||
from core.inference.mlx_inference import _kv_prefix_coverage, _MLXPromptCacheHistory
|
||||
|
||||
def feed(entry, n):
|
||||
for _ in range(n):
|
||||
block = mx.zeros((1, 2, 1, 4), dtype = mx.float16)
|
||||
entry.update_and_fetch(block, block)
|
||||
mx.eval(entry.state)
|
||||
return entry
|
||||
|
||||
plain = feed(KVCache(), 30)
|
||||
unwrapped = feed(RotatingKVCache(max_size = 100, keep = 2), 30)
|
||||
wrapped = feed(RotatingKVCache(max_size = 10, keep = 2), 30)
|
||||
chunked = feed(ChunkedKVCache(chunk_size = 8), 30)
|
||||
slid = feed(ChunkedKVCache(chunk_size = 8), 30)
|
||||
slid.maybe_trim_front()
|
||||
|
||||
assert _kv_prefix_coverage([plain]) == 30
|
||||
assert _kv_prefix_coverage([unwrapped]) == 30
|
||||
assert _kv_prefix_coverage([chunked]) == 30
|
||||
assert wrapped.offset == 30 and wrapped.state[0].shape[2] == 10
|
||||
assert _kv_prefix_coverage([wrapped]) is None
|
||||
assert slid.start_position > 0
|
||||
assert _kv_prefix_coverage([slid]) is None
|
||||
assert _kv_prefix_coverage([CacheList(feed(KVCache(), 30), feed(KVCache(), 30))]) == 30
|
||||
assert _kv_prefix_coverage([CacheList(feed(KVCache(), 30), wrapped)]) is None
|
||||
assert _kv_prefix_coverage([feed(KVCache(), 30), feed(KVCache(), 29)]) is None
|
||||
assert _kv_prefix_coverage([]) is None
|
||||
|
||||
history = _MLXPromptCacheHistory(6, 1 << 40)
|
||||
for unsafe in (wrapped, slid):
|
||||
history.insert("key", list(range(30)), [unsafe])
|
||||
assert "key" not in history._lru.entries
|
||||
|
||||
history.insert("key", list(range(30)), [plain])
|
||||
assert tuple(range(30)) in history._lru.entries["key"]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue