diff --git a/inference_engine/distributed/cache_budget.py b/inference_engine/distributed/cache_budget.py new file mode 100644 index 00000000..84300c0e --- /dev/null +++ b/inference_engine/distributed/cache_budget.py @@ -0,0 +1,15 @@ +"""Platform-neutral cache-budget policy for model-loaded offload workers.""" + + +def adaptive_cache_budget( + *, + total_bytes: int, + active_model_bytes: int, + ceiling_bytes: int, + minimum_bytes: int, + reserve_bytes: int, +) -> int: + if min(total_bytes, ceiling_bytes, minimum_bytes) <= 0 or reserve_bytes < 0: + raise ValueError("adaptive cache budget inputs are invalid") + available = max(0, total_bytes - active_model_bytes - reserve_bytes) + return max(minimum_bytes, min(ceiling_bytes, available)) diff --git a/scripts/start_prefill_worker_node.py b/scripts/start_prefill_worker_node.py index e658719f..37462a8c 100644 --- a/scripts/start_prefill_worker_node.py +++ b/scripts/start_prefill_worker_node.py @@ -33,6 +33,7 @@ NodeEndpoint, PrefillWorkerCapability, ) +from inference_engine.distributed.cache_budget import adaptive_cache_budget from inference_engine.distributed.exchange import ( add_capability_service, exchange_once, @@ -68,20 +69,6 @@ def physical_memory_bytes() -> int: return 0 -def adaptive_cache_budget( - *, - total_bytes: int, - active_model_bytes: int, - ceiling_bytes: int, - minimum_bytes: int, - reserve_bytes: int, -) -> int: - if min(total_bytes, ceiling_bytes, minimum_bytes) <= 0 or reserve_bytes < 0: - raise ValueError("adaptive cache budget inputs are invalid") - available = max(0, total_bytes - active_model_bytes - reserve_bytes) - return max(minimum_bytes, min(ceiling_bytes, available)) - - def mlx_active_memory_bytes() -> int: try: import mlx.core as mx diff --git a/tests/inference_engine/bridge/test_prefill_worker_memory.py b/tests/inference_engine/bridge/test_prefill_worker_memory.py index 767a62be..357484c0 100644 --- a/tests/inference_engine/bridge/test_prefill_worker_memory.py +++ b/tests/inference_engine/bridge/test_prefill_worker_memory.py @@ -1,6 +1,6 @@ import pytest -from scripts.start_prefill_worker_node import adaptive_cache_budget +from inference_engine.distributed.cache_budget import adaptive_cache_budget def test_adaptive_budget_uses_only_model_headroom():