Handle low memory better

This commit is contained in:
Ryuichi Leo Takashige
2026-02-25 19:25:18 +00:00
parent ba611f9cd0
commit 2f719d62a7
6 changed files with 337 additions and 8 deletions
+2 -1
View File
@@ -249,7 +249,8 @@ class ChunkedKVCache(KVCache):
...
class CacheList(_BaseCache):
def __init__(self, *caches) -> None: ...
caches: tuple[_BaseCache, ...]
def __init__(self, *caches: _BaseCache) -> None: ...
def __getitem__(self, idx): ...
def is_trimmable(self): # -> bool:
...
+181
View File
@@ -0,0 +1,181 @@
#!/usr/bin/env python3
# Test OOM prevention by sending increasingly large contexts to exo.
#
# Usage:
# 1. Start exo with a model loaded (in another terminal):
# uv run exo
#
# 2. Run this script:
# uv run python scripts/test_oom_prevention.py
#
# 3. To force the OOM check to trigger sooner, lower the threshold:
# EXO_MEMORY_THRESHOLD=0.5 uv run exo
# Then run this script -- it should trigger much earlier.
import json
import sys
import time
import httpx
BASE_URL = "http://localhost:52415"
MODEL = "mlx-community/Llama-3.3-70B-Instruct-4bit"
def send_chat(
messages: list[dict[str, str]], max_tokens: int = 100, stream: bool = True
) -> str | None:
payload = {
"model": MODEL,
"messages": messages,
"max_tokens": max_tokens,
"stream": stream,
}
if stream:
return _send_streaming(payload)
return _send_non_streaming(payload)
def _send_streaming(payload: dict[str, object]) -> str | None:
collected = ""
error = None
with httpx.stream(
"POST", f"{BASE_URL}/v1/chat/completions", json=payload, timeout=120
) as resp:
for line in resp.iter_lines():
if not line.startswith("data: "):
continue
data = line[len("data: ") :]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
if "error" in chunk:
err = chunk["error"]
error = (
err.get("message", str(err))
if isinstance(err, dict)
else str(err)
)
break
delta = chunk.get("choices", [{}])[0].get("delta", {})
collected += delta.get("content", "")
except json.JSONDecodeError:
pass
if error:
return f"ERROR: {error}"
return collected
def _send_non_streaming(payload: dict[str, object]) -> str | None:
resp = httpx.post(f"{BASE_URL}/v1/chat/completions", json=payload, timeout=120)
if resp.status_code != 200:
return f"ERROR (HTTP {resp.status_code}): {resp.text}"
data = resp.json()
return data.get("choices", [{}])[0].get("message", {}).get("content", "")
def make_large_message(token_count_approx: int) -> str:
return "buffalo " * token_count_approx
def test_oom_prevention() -> None:
print("=" * 60)
print("OOM Prevention Test")
print("=" * 60)
try:
resp = httpx.get(f"{BASE_URL}/v1/models", timeout=5)
models = resp.json()
print(
f"\nConnected to exo. Models: {[m['id'] for m in models.get('data', [])]}"
)
except Exception as e:
print(f"\nCannot connect to exo at {BASE_URL}: {e}")
print("Start exo first: uv run exo")
sys.exit(1)
print("\n--- Test 1: Small request (should succeed) ---")
result = send_chat(
[{"role": "user", "content": "Say hello in exactly 5 words."}],
max_tokens=50,
)
print(f"Response: {result}")
if result and result.startswith("ERROR"):
print("FAIL: Small request should not fail")
sys.exit(1)
print("PASS")
sizes = [1000, 5000, 10000, 20000, 30000, 50000, 80000, 120000]
max_gen = 8000
print("\n--- Test 2: Escalating context sizes ---")
print(f"{'Size':>10} | {'Max Gen':>8} | {'Result':>10} | Details")
print("-" * 70)
for size in sizes:
large_msg = make_large_message(size)
messages = [
{
"role": "user",
"content": f"Summarize this in one sentence:\n{large_msg}",
},
]
t0 = time.time()
try:
result = send_chat(messages, max_tokens=max_gen)
elapsed = time.time() - t0
except Exception as e:
print(f"{size:>10} | {max_gen:>8} | {'EXCEPTION':>10} | {e}")
continue
if result is None:
print(f"{size:>10} | {max_gen:>8} | {'NONE':>10} | No response")
elif "not enough memory" in result.lower() or "ERROR" in result:
print(f"{size:>10} | {max_gen:>8} | {'OOM CAUGHT':>10} | {result[:80]}...")
print(
f"\nOOM prevention triggered at ~{size} prompt tokens "
f"+ {max_gen} gen tokens"
)
print("SUCCESS: The system caught the OOM before crashing!")
return
else:
preview = result[:60].replace("\n", " ")
print(
f"{size:>10} | {max_gen:>8} | {'OK':>10} | "
f"{preview}... ({elapsed:.1f}s)"
)
print("\nNOTE: OOM prevention did not trigger at any tested size.")
print("Either your machine has enough memory, or try:")
print(" EXO_MEMORY_THRESHOLD=0.5 uv run exo")
def test_with_low_threshold() -> None:
# Requires: EXO_MEMORY_THRESHOLD=0.3 uv run exo
print("\n--- Test 3: Forced OOM (requires EXO_MEMORY_THRESHOLD=0.3) ---")
result = send_chat(
[
{
"role": "user",
"content": "Write a long essay about the history of computing.",
}
],
max_tokens=16000,
)
if result and ("not enough memory" in result.lower() or "ERROR" in result):
print(f"PASS: OOM prevention triggered: {result[:100]}...")
else:
preview = (result or "")[:100]
print(
f"Did not trigger. Is EXO_MEMORY_THRESHOLD set low enough? Got: {preview}"
)
if __name__ == "__main__":
test_oom_prevention()
if "--force" in sys.argv:
test_with_low_threshold()
+52 -2
View File
@@ -32,7 +32,7 @@ def _default_memory_threshold() -> float:
return 0.70
_MEMORY_THRESHOLD = float(
MEMORY_THRESHOLD = float(
os.environ.get("EXO_MEMORY_THRESHOLD", _default_memory_threshold())
)
@@ -92,6 +92,15 @@ class KVPrefixCache:
self._snapshots.clear()
self._last_used.clear()
def force_evict_all(self) -> int:
count = len(self.caches)
self.clear()
if count > 0:
logger.info(
f"Force-evicted all {count} prefix cache entries due to memory pressure"
)
return count
def add_kv_cache(
self,
prompt_tokens: mx.array,
@@ -217,7 +226,7 @@ class KVPrefixCache:
# Evict LRU entries until below threshold
while (
len(self.caches) > 0
and self.get_memory_used_percentage() > _MEMORY_THRESHOLD
and self.get_memory_used_percentage() > MEMORY_THRESHOLD
):
lru_index = self._last_used.index(min(self._last_used))
evicted_tokens = len(self.prompts[lru_index])
@@ -310,6 +319,47 @@ def get_memory_used_percentage() -> float:
return float(mem.percent / 100)
def _measure_single_cache_bytes(
entry: KVCache | RotatingKVCache | QuantizedKVCache | ArraysCache | CacheList,
) -> int:
if isinstance(entry, CacheList):
return sum(
_measure_single_cache_bytes(c) # pyright: ignore[reportArgumentType]
for c in entry.caches
)
total = 0
for attr_name in ("keys", "values"):
val: object = getattr(entry, attr_name, None)
if val is None:
continue
if isinstance(val, mx.array):
total += val.nbytes
elif isinstance(val, (tuple, list)):
# QuantizedKVCache stores tuples of arrays (data, scales, biases)
for arr in val: # pyright: ignore[reportUnknownVariableType]
if isinstance(arr, mx.array):
total += arr.nbytes
if isinstance(entry, ArraysCache):
state = entry.state # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
for arr in state: # pyright: ignore[reportUnknownVariableType]
if isinstance(arr, mx.array):
total += arr.nbytes
return total
def measure_cache_bytes(cache: KVCacheType) -> int:
return sum(_measure_single_cache_bytes(c) for c in cache)
def measure_kv_cache_bytes_per_token(cache: KVCacheType) -> int:
offset = cache_length(cache)
if offset == 0:
return 0
return measure_cache_bytes(cache) // offset
def make_kv_cache(
model: Model, max_kv_size: int | None = None, keep: int = 0
) -> KVCacheType:
@@ -4,6 +4,7 @@ from copy import deepcopy
from typing import Callable, Generator, cast, get_args
import mlx.core as mx
import psutil
from mlx_lm.generate import stream_generate
from mlx_lm.models.cache import ArraysCache, RotatingKVCache
from mlx_lm.sample_utils import make_sampler
@@ -27,11 +28,14 @@ from exo.shared.types.worker.runner_response import (
from exo.worker.engines.mlx import Model
from exo.worker.engines.mlx.auto_parallel import set_pipeline_prefill
from exo.worker.engines.mlx.cache import (
MEMORY_THRESHOLD,
CacheSnapshot,
KVPrefixCache,
encode_prompt,
get_memory_used_percentage,
has_non_kv_caches,
make_kv_cache,
measure_kv_cache_bytes_per_token,
snapshot_ssm_states,
)
from exo.worker.engines.mlx.constants import (
@@ -148,7 +152,13 @@ def warmup_inference(
model: Model,
tokenizer: TokenizerWrapper,
group: mx.distributed.Group | None,
) -> int:
) -> tuple[int, int]:
"""Run warmup inference and measure KV cache cost per token.
Returns:
(tokens_generated, bytes_per_token) where bytes_per_token is the
measured KV cache memory consumed per token across all local layers.
"""
content = "Prompt to warm up the inference engine. Repeat this."
warmup_prompt = apply_chat_template(
@@ -187,9 +197,13 @@ def warmup_inference(
logger.info("Generated ALL warmup tokens")
# Measure KV cache bytes per token from the populated warmup cache
bytes_per_token = measure_kv_cache_bytes_per_token(cache)
logger.info(f"Measured KV cache cost: {bytes_per_token} bytes per token")
mx_barrier(group)
return tokens_generated
return tokens_generated, bytes_per_token
def ban_token_ids(token_ids: list[int]) -> Callable[[mx.array, mx.array], mx.array]:
@@ -267,6 +281,65 @@ def extract_top_logprobs(
return selected_logprob, top_logprob_items
def _check_memory_budget(
bytes_per_token: int,
total_sequence_tokens: int,
kv_prefix_cache: KVPrefixCache | None,
) -> str | None:
"""Check if enough memory is available for the estimated KV cache.
Uses the same memory pressure system as prefix cache eviction.
If memory would exceed the threshold, tries evicting prefix caches first.
Returns None if OK, or an error message string if OOM is predicted.
"""
if bytes_per_token == 0:
return None
total_ram = psutil.virtual_memory().total
estimated_cache_bytes = bytes_per_token * total_sequence_tokens
current_pressure = (
kv_prefix_cache.get_memory_used_percentage()
if kv_prefix_cache is not None
else get_memory_used_percentage()
)
projected_pressure = current_pressure + (estimated_cache_bytes / total_ram)
logger.info(
f"Memory check: {total_sequence_tokens} tokens × {bytes_per_token} B/tok "
f"= {estimated_cache_bytes / (1024**2):.1f} MB, "
f"pressure {current_pressure:.1%} → projected {projected_pressure:.1%} "
f"(threshold {MEMORY_THRESHOLD:.1%})"
)
if projected_pressure <= MEMORY_THRESHOLD:
return None
# Try evicting all prefix caches
if kv_prefix_cache is not None:
evicted = kv_prefix_cache.force_evict_all()
if evicted > 0:
mx.clear_cache()
current_pressure = kv_prefix_cache.get_memory_used_percentage()
projected_pressure = current_pressure + (estimated_cache_bytes / total_ram)
logger.info(
f"After evicting {evicted} prefix cache entries: "
f"pressure {current_pressure:.1%} → projected {projected_pressure:.1%}"
)
if projected_pressure <= MEMORY_THRESHOLD:
return None
needed_mb = estimated_cache_bytes / (1024**2)
headroom_mb = max(0, (MEMORY_THRESHOLD - current_pressure) * total_ram) / (1024**2)
return (
f"Not enough memory for this conversation. "
f"Estimated KV cache need: {needed_mb:.0f} MB, "
f"available headroom: {headroom_mb:.0f} MB. "
f"Please start a new conversation or compact your messages to continue."
)
def mlx_generate(
model: Model,
tokenizer: TokenizerWrapper,
@@ -275,6 +348,7 @@ def mlx_generate(
kv_prefix_cache: KVPrefixCache | None,
group: mx.distributed.Group | None,
on_prefill_progress: Callable[[int, int], None] | None = None,
bytes_per_token: int = 0,
) -> Generator[GenerationResponse]:
# Ensure that generation stats only contains peak memory for this generation
mx.reset_peak_memory()
@@ -307,6 +381,25 @@ def mlx_generate(
f"KV cache hit: {prefix_hit_length}/{len(all_prompt_tokens)} tokens cached ({100 * prefix_hit_length / len(all_prompt_tokens):.1f}%)"
)
# OOM prevention: check if the full sequence will fit in memory
if bytes_per_token > 0:
max_tokens = task.max_output_tokens or MAX_TOKENS
total_sequence_tokens = len(all_prompt_tokens) + max_tokens
oom_error = _check_memory_budget(
bytes_per_token=bytes_per_token,
total_sequence_tokens=total_sequence_tokens,
kv_prefix_cache=kv_prefix_cache,
)
if oom_error is not None:
logger.warning(f"OOM prevention triggered: {oom_error}")
yield GenerationResponse(
text=oom_error,
token=0,
finish_reason="error",
usage=None,
)
return
logits_processors: list[Callable[[mx.array, mx.array], mx.array]] = []
if is_bench:
# Only sample length eos tokens
@@ -114,6 +114,7 @@ def main(
group = None
kv_prefix_cache: KVPrefixCache | None = None
check_for_cancel_every: int | None = None
bytes_per_token: int = 0
current_status: RunnerStatus = RunnerIdle()
logger.info("runner created")
@@ -225,12 +226,14 @@ def main(
assert tokenizer
t = time.monotonic()
toks = warmup_inference(
toks, bytes_per_token = warmup_inference(
model=cast(Model, inference_model),
tokenizer=tokenizer,
group=group,
)
logger.info(f"warmed up by generating {toks} tokens")
logger.info(
f"warmed up by generating {toks} tokens, {bytes_per_token} bytes/token for KV cache"
)
check_for_cancel_every = min(
math.ceil(toks / min(time.monotonic() - t, 0.001)), 100
)
@@ -310,6 +313,7 @@ def main(
kv_prefix_cache=kv_prefix_cache,
on_prefill_progress=on_prefill_progress,
group=group,
bytes_per_token=bytes_per_token,
)
if tokenizer.has_thinking:
@@ -114,7 +114,7 @@ def patch_out_mlx(monkeypatch: pytest.MonkeyPatch):
# initialize_mlx returns a mock group
monkeypatch.setattr(mlx_runner, "initialize_mlx", make_nothin(MockGroup()))
monkeypatch.setattr(mlx_runner, "load_mlx_items", make_nothin((1, MockTokenizer)))
monkeypatch.setattr(mlx_runner, "warmup_inference", make_nothin(1))
monkeypatch.setattr(mlx_runner, "warmup_inference", make_nothin((1, 0)))
monkeypatch.setattr(mlx_runner, "_check_for_debug_prompts", nothin)
monkeypatch.setattr(mlx_runner, "mx_any", make_nothin(False))
# Mock apply_chat_template since we're using a fake tokenizer (integer 1).