Handle low memory better
This commit is contained in:
@@ -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:
|
||||
...
|
||||
|
||||
@@ -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()
|
||||
@@ -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).
|
||||
|
||||
Reference in New Issue
Block a user