Skip to content
Open
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -1218,6 +1218,7 @@ Structured outputs are not currently supported with speculative decoding.
- `/v1/realtime` - WebSocket-based realtime full-duplex speech. See the [Nemotron VoiceChat guide](mlx_vlm/models/nemotron_voicechat/README.md#realtime-websocket-api) for supported models and usage.
- `/health` - Check server status
- `/metrics` and `/v1/metrics` - Inspect rolling request metrics, throughput, and runtime counters
- `/cache/offload` and `/v1/cache/offload` - Write a conversation's prefix cache to disk and free the memory it holds, keeping the prefix reusable
- `/unload` - Unload all loaded model caches from memory

#### Usage Examples
Expand Down
107 changes: 107 additions & 0 deletions mlx_vlm/apc.py
Original file line number Diff line number Diff line change
Expand Up @@ -3087,6 +3087,7 @@ def __init__(
self._free_push(b)
self.hash_table: dict[int, APCBlock] = {}
self._exact_cache: "OrderedDict[int, APCExactCacheEntry]" = OrderedDict()
self._seen_extra_hashes: "OrderedDict[int, None]" = OrderedDict()
self.stats = APCStats()
self.lock = threading.RLock()
self.disk = disk
Expand Down Expand Up @@ -3902,6 +3903,9 @@ def store_kv_blocks(
self.stats.stores += 1
self.stats.served_tokens += self.block_size
parent = h
self._seen_extra_hashes[int(extra_hash)] = None
while len(self._seen_extra_hashes) > 16:
self._seen_extra_hashes.popitem(last=False)
if self.disk is not None and disk_blocks:
try:
self.disk.save_layer_major_blocks(
Expand All @@ -3912,6 +3916,109 @@ def store_kv_blocks(
self.stats.pool_used = sum(1 for x in self.pool if x.block_hash is not None)
return new_blocks

def _offload_candidate_salts(self, extra_hash: Optional[int]) -> List[int]:
salts = [] if extra_hash is None else [int(extra_hash)]
for salt in self._seen_extra_hashes:
if salt not in salts:
salts.append(salt)
return salts or [0]

def offload_prefix(
self, token_ids: Sequence[int], extra_hash: Optional[int] = None
) -> dict:
"""Persist a prefix's blocks to disk, then free the memory they hold.

Blocks another request still holds, and blocks that are not on disk,
stay resident: releasing either would drop state a caller still needs.
"""
if self.disk is not None:
self.disk.flush()
with self.lock:
resident: List[Tuple[int, APCBlock]] = []
for salt in self._offload_candidate_salts(extra_hash):
walked: List[Tuple[int, APCBlock]] = []
parent = SEED_PARENT_HASH
for i in range(len(token_ids) // self.block_size):
chunk = tuple(
int(t)
for t in token_ids[
i * self.block_size : (i + 1) * self.block_size
]
)
h = _hash_tokens(parent, chunk, salt)
block = self.hash_table.get(h)
if block is None or block.token_ids != chunk:
break
walked.append((h, block))
parent = h
if len(walked) > len(resident):
resident = walked

released = in_use = unpersisted = freed_bytes = 0
for h, block in resident:
if block.ref_cnt > 0:
in_use += 1
continue
if self.disk is None or not self.disk.has(h):
unpersisted += 1
continue
freed_bytes += block.resident_bytes()
self._free_remove(block)
if self.hash_table.get(h) is block:
del self.hash_table[h]
self.stats.evictions += 1
block.block_hash = None
block.token_ids = ()
block.release_components()
self._free_push(block)
released += 1

self.stats.pool_used = sum(1 for x in self.pool if x.block_hash is not None)

# Custom cache layouts are kept as whole-prefix snapshots rather than
# blocks, so they have to be spilled and dropped on their own terms.
prefix = tuple(int(t) for t in token_ids)
snapshots = released_snapshots = retained_snapshots = 0
for key, entry in list(self._exact_cache.items()):
if prefix[: len(entry.token_ids)] != entry.token_ids:
continue
snapshots += 1
persisted = self.disk is not None and (
self.disk.find_exact_prefix(
entry.token_ids,
extra_hash=entry.extra_hash,
block_size=self.block_size,
)
is not None
)
if not persisted and self.disk is not None:
persisted = self.disk.save_exact_cache(
key,
entry.token_ids,
entry.extra_hash,
entry.prompt_cache,
synchronous=True,
)
if not persisted:
retained_snapshots += 1
continue
with self.lock:
dropped = self._exact_cache.pop(key, None)
if dropped is not None:
freed_bytes += _cache_nbytes(dropped.prompt_cache)
released_snapshots += 1

return {
"matched_blocks": len(resident),
"released_blocks": released,
"released_tokens": released * self.block_size,
"retained_in_use": in_use,
"retained_unpersisted": unpersisted + retained_snapshots,
"matched_snapshots": snapshots,
"released_snapshots": released_snapshots,
"freed_bytes": freed_bytes,
}

def stats_snapshot(self) -> dict:
with self.lock:
self.stats.pool_used = sum(1 for x in self.pool if x.block_hash is not None)
Expand Down
89 changes: 89 additions & 0 deletions mlx_vlm/server/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -333,6 +333,8 @@ def register_routes(app, deps):
app.get("/v1/responses/{response_id}/input_items", include_in_schema=False)(
responses_input_items_endpoint
)
app.post("/cache/offload")(cache_offload_endpoint)
app.post("/v1/cache/offload", include_in_schema=False)(cache_offload_endpoint)
app.post("/responses")(responses_endpoint)
app.post("/v1/responses", include_in_schema=False)(responses_endpoint)
app.post("/chat/completions", response_model=None)(chat_completions_endpoint)
Expand Down Expand Up @@ -756,6 +758,93 @@ async def responses_input_tokens_endpoint(request: Request):
raise HTTPException(status_code=400, detail=str(e))


def _conversation_prefix_ids(processor, config, prompt, images) -> list:
"""The token ids generation ran over for this conversation.

Re-tokenizing the rendered prompt is not enough: a template that emits its
own BOS shifts every id by one and misses the cache entirely.
"""
images = images or None
try:
if runtime.response_generator is not None:
raw_inputs = runtime.response_generator._cpu_preprocess(
prompt, images, None
)
else:
raw_inputs = prepare_inputs(
processor,
images=images,
prompts=prompt,
image_token_index=getattr(config, "image_token_index", None),
)
input_ids = raw_inputs["input_ids"]
except Exception as e:
logger.debug("Could not derive conversation prefix ids: %s", e)
return []
ids = input_ids[0] if getattr(input_ids, "ndim", 1) > 1 else input_ids
return [int(t) for t in ids]


async def cache_offload_endpoint(request: Request):
"""Write a conversation's prefix cache to disk and free the memory it holds.

The conversation is supplied the way a Responses request supplies it, so
the prefix released is the one generation actually ran over.
"""
body = await request.json()
openai_request = OpenAIRequest(**body)

try:
if openai_request.input is None:
raise HTTPException(status_code=400, detail="Missing input.")

prompt_items = _response_chain_items(
openai_request.previous_response_id
) + _normalize_response_input(openai_request.input)
chat_messages, images = _response_items_to_chat(prompt_items)
_normalize_response_instruction_messages(
chat_messages, openai_request.instructions
)
_ensure_effective_input(chat_messages, images=images)

if runtime.apc_manager is None:
raise HTTPException(status_code=409, detail="Prompt cache is disabled.")

model, processor, config = get_cached_model(
openai_request.model, _adapter_path_or_inherit(openai_request)
)
del model
gen_args = _build_gen_args(
openai_request, processor, tenant_id=_read_tenant_id(request)
)
formatted_prompt = apply_chat_template(
processor,
config,
chat_messages,
num_images=len(images),
**gen_args.to_template_kwargs(),
)
token_ids = _conversation_prefix_ids(
processor, config, formatted_prompt, images
)
if not token_ids:
raise HTTPException(
status_code=422, detail="Could not resolve the conversation prefix."
)

result = runtime.apc_manager.offload_prefix(token_ids)
mx.clear_cache()
gc.collect()
return result
except HTTPException:
raise
except Exception as e:
logger.exception("Unexpected error in /cache/offload endpoint: %s", e)
raise HTTPException(
status_code=500, detail=f"An unexpected error occurred: {e}"
)


async def responses_retrieve_endpoint(response_id: str):
with response_store_lock:
stored = response_store.get(response_id)
Expand Down
144 changes: 144 additions & 0 deletions mlx_vlm/tests/test_models.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import importlib
import inspect
import math
import pathlib
import tempfile
import threading
import unittest
from types import SimpleNamespace
Expand Down Expand Up @@ -20108,3 +20110,145 @@ def test_sanitize_adds_language_model_prefix(self):
weights = {"model.embedding.weight": mx.zeros((128, 64))}
sanitized = model.sanitize(weights)
self.assertIn("language_model.model.embedding.weight", sanitized)


import pathlib
import tempfile
from unittest.mock import patch


class TestAPCOffloadPrefix(unittest.TestCase):
"""Offload must persist before it frees, and never drop live state."""

def _manager(self, root):
from mlx_vlm.apc import APCManager, DiskBlockStore

return APCManager(
num_blocks=64,
block_size=16,
disk=DiskBlockStore(pathlib.Path(root), "offload-test"),
)

def _store(self, manager, ids):
keys = mx.random.normal((1, 1, len(ids), 4))
values = mx.random.normal((1, 1, len(ids), 4))
blocks = manager.store_kv_blocks(ids, [keys], [values])
manager.release(blocks)
return blocks

def test_offload_frees_memory_and_leaves_the_prefix_on_disk(self):
with tempfile.TemporaryDirectory() as root:
manager = self._manager(root)
ids = list(range(256))
self._store(manager, ids)

before = manager.stats_snapshot()["pool_used"]
self.assertGreater(before, 0)
self.assertGreater(manager.resident_bytes(), 0)

result = manager.offload_prefix(ids)

self.assertEqual(result["matched_blocks"], 16)
self.assertEqual(result["released_blocks"], 16)
self.assertEqual(result["released_tokens"], 256)
self.assertEqual(result["retained_in_use"], 0)
self.assertEqual(result["retained_unpersisted"], 0)
self.assertGreater(result["freed_bytes"], 0)
self.assertEqual(manager.stats_snapshot()["pool_used"], 0)

# the prefix survives on disk for a later matching request
self.assertTrue(all(manager.disk.has(h) for h in manager.disk._index))
self.assertEqual(manager.lookup_prefix(ids)[1], 0)
self.assertGreater(manager.lookup_prefix_disk_cache(ids)[1], 0)
manager.close()

def test_offload_keeps_blocks_another_request_still_holds(self):
with tempfile.TemporaryDirectory() as root:
manager = self._manager(root)
ids = list(range(256))
self._store(manager, ids)

held, _ = manager.lookup_prefix(ids)
self.assertGreater(len(held), 0)

result = manager.offload_prefix(ids)
self.assertEqual(result["released_blocks"], 0)
self.assertEqual(result["retained_in_use"], len(held))
self.assertEqual(manager.stats_snapshot()["pool_used"], len(held))

manager.release(held)
after = manager.offload_prefix(ids)
self.assertEqual(after["released_blocks"], len(held))
self.assertEqual(manager.stats_snapshot()["pool_used"], 0)
manager.close()

def test_offload_retains_blocks_that_never_reached_disk(self):
with tempfile.TemporaryDirectory() as root:
manager = self._manager(root)
ids = list(range(256))
self._store(manager, ids)

with patch.object(manager.disk, "has", return_value=False):
result = manager.offload_prefix(ids)

self.assertEqual(result["released_blocks"], 0)
self.assertEqual(result["retained_unpersisted"], 16)
self.assertEqual(manager.stats_snapshot()["pool_used"], 16)
self.assertGreater(manager.lookup_prefix(ids)[1], 0)
manager.close()

def test_offload_is_a_no_op_for_a_prefix_that_was_never_cached(self):
with tempfile.TemporaryDirectory() as root:
manager = self._manager(root)
self._store(manager, list(range(256)))

result = manager.offload_prefix(list(range(1000, 1256)))

self.assertEqual(result["matched_blocks"], 0)
self.assertEqual(result["released_blocks"], 0)
self.assertEqual(manager.stats_snapshot()["pool_used"], 16)
manager.close()

def test_offloaded_blocks_return_to_the_pool_without_leaking(self):
with tempfile.TemporaryDirectory() as root:
manager = self._manager(root)
for start in range(0, 5):
ids = list(range(start * 1000, start * 1000 + 256))
self._store(manager, ids)
manager.offload_prefix(ids)

self.assertEqual(manager.stats_snapshot()["pool_used"], 0)
self.assertEqual(sum(b.ref_cnt for b in manager.pool), 0)
free = 0
node = manager._free_head
while node is not None:
free += 1
node = node.next
self.assertEqual(free, manager.num_blocks)
manager.close()


class TestAPCOffloadSaltDiscovery(unittest.TestCase):
"""The salt folds in model and processor, so offload cannot assume zero."""

def test_offload_finds_a_prefix_stored_under_a_model_specific_salt(self):
from mlx_vlm.apc import APCManager, DiskBlockStore

with tempfile.TemporaryDirectory() as root:
manager = APCManager(
num_blocks=64,
block_size=16,
disk=DiskBlockStore(pathlib.Path(root), "salt-test"),
)
ids = list(range(256))
keys = mx.random.normal((1, 1, len(ids), 4))
manager.release(
manager.store_kv_blocks(ids, [keys], [keys], extra_hash=987654321)
)

result = manager.offload_prefix(ids)

self.assertEqual(result["matched_blocks"], 16)
self.assertEqual(result["released_blocks"], 16)
self.assertEqual(manager.stats_snapshot()["pool_used"], 0)
manager.close()
Loading
Loading